結果

問題 No.3762 Glowing Utility Pole
コンテスト
ユーザー ei13333333
提出日時 2026-10-04 18:39:03
言語 C++23(gcc16)
(gcc 16.1.0 + boost 1.92.0 + ACL)
コンパイル:
g++-16 -O2 -lm -std=c++23 -Wuninitialized -DONLINE_JUDGE -o a.out _filename_
実行:
./a.out
結果
TLE  
実行時間 -
コード長 7,224 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 2,205 ms
コンパイル使用メモリ 366,208 KB
実行使用メモリ 10,496 KB
最終ジャッジ日時 2026-10-09 20:52:30
合計ジャッジ時間 13,740 ms
ジャッジサーバーID
(参考情報)
judge1_1 / judge5_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 3
other AC * 46 TLE * 1
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

#include <bits/stdc++.h>
using namespace std;

#define FOR(i, n, m) for (int i = n; i < (int)m; ++i)
#define REP(i, n) FOR(i, 0, n)
#define REP_1(i, n) for (int i = 1; i <= (int)n; ++i)
#define RFOR(i, n, m) for (int i = (int)n - 1; i >= (int)m; --i)
#define RREP(i, n) RFOR(i, n, 0)
#define RREP_1(i, n) for (int i = (int)n; i >= 1; --i)

#define ALL(v) v.begin(), v.end()
#define RALL(v) v.rbegin(), v.rend()
#define SIZE(v) (int)v.size()
#define EMPTY(v) v.empty()
#define SORT(v) sort(ALL(v))
#define RSORT(v) sort(RALL(v))
#define REVERSE(v) reverse(ALL(v))
#define UNIQUE(v) (SORT(v), v.erase(unique(ALL(v)), v.end()))

#define PB push_back
#define EB emplace_back
#define MP make_pair

#define YES() cout << "YES\n"
#define NO() cout << "NO\n"
#define Yes() cout << "Yes\n"
#define No() cout << "No\n"
#define YESNO(cond) cout << ((cond) ? "YES" : "NO") << '\n'
#define YesNo(cond) cout << ((cond) ? "Yes" : "No") << '\n'

#define IN(x, a, b) ((a) <= (x) && (x) < (b))
#define BETWEEN(x, a, b) ((a) <= (x) && (x) <= (b))

#define FASTIO()                                                               \
    ios::sync_with_stdio(false), cin.tie(nullptr), cout.tie(nullptr)
#define PRECISION(n) cout << fixed << setprecision(n)

using P = pair<int, int>;
using ll = long long;
using ull = unsigned long long;
using ld = long double;

template <class T> using min_queue = priority_queue<T, vector<T>, greater<T>>;
template <class T> using max_queue = priority_queue<T>;

constexpr ll INF = 1000000000;
constexpr ll INFL = (ll)1000000000000001000LL;
constexpr ll MOD = 998244353;
constexpr ld PI = 3.141592653589793238462643383279;
constexpr ld EPS = 1e-9;

constexpr int dx4[] = {-1, 1, 0, 0};
constexpr int dy4[] = {0, 0, -1, 1};
constexpr int dx8[] = {0, 1, 1, 1, 0, -1, -1, -1};
constexpr int dy8[] = {1, 1, 0, -1, -1, -1, 0, 1};

struct modint {
    ll n;

  public:
    modint() : n(0) {}
    modint(ll x) : n(((x % MOD) + MOD) % MOD) {}

    ll val() const { return n; }

    modint pow(ll m) const {
        modint r = 1, a = *this;
        while (m > 0) {
            if (m & 1) r *= a;
            a *= a;
            m >>= 1;
        }
        return r;
    }

    modint inv() const { return pow(MOD - 2); }

    modint &operator++() { return *this += 1; }
    modint &operator--() { return *this -= 1; }
    modint operator++(int) {
        modint ret = *this;
        ++*this;
        return ret;
    }
    modint operator--(int) {
        modint ret = *this;
        --*this;
        return ret;
    }

    modint operator+() const { return *this; }
    modint operator-() const { return modint() - *this; }

    friend bool operator==(const modint &lhs, const modint &rhs) {
        return lhs.n == rhs.n;
    }
    friend bool operator!=(const modint &lhs, const modint &rhs) {
        return lhs.n != rhs.n;
    }
    friend bool operator<(const modint &lhs, const modint &rhs) {
        return lhs.n < rhs.n;
    }
    friend bool operator<=(const modint &lhs, const modint &rhs) {
        return lhs.n <= rhs.n;
    }
    friend bool operator>(const modint &lhs, const modint &rhs) {
        return lhs.n > rhs.n;
    }
    friend bool operator>=(const modint &lhs, const modint &rhs) {
        return lhs.n >= rhs.n;
    }

    friend modint &operator+=(modint &lhs, const modint &rhs) {
        lhs.n += rhs.n;
        if (lhs.n >= MOD) lhs.n -= MOD;
        return lhs;
    }
    friend modint &operator-=(modint &lhs, const modint &rhs) {
        lhs.n -= rhs.n;
        if (lhs.n < 0) lhs.n += MOD;
        return lhs;
    }
    friend modint &operator*=(modint &lhs, const modint &rhs) {
        lhs.n = (lhs.n * rhs.n) % MOD;
        return lhs;
    }
    friend modint &operator/=(modint &lhs, const modint &rhs) {
        return lhs *= rhs.inv();
    }

    friend modint operator+(const modint &lhs, const modint &rhs) {
        modint res = lhs;
        res += rhs;
        return res;
    }
    friend modint operator-(const modint &lhs, const modint &rhs) {
        modint res = lhs;
        res -= rhs;
        return res;
    }
    friend modint operator*(const modint &lhs, const modint &rhs) {
        modint res = lhs;
        res *= rhs;
        return res;
    }
    friend modint operator/(const modint &lhs, const modint &rhs) {
        modint res = lhs;
        res /= rhs;
        return res;
    }

    friend istream &operator>>(istream &is, modint &m) {
        ll x;
        is >> x;
        m = modint(x);
        return is;
    }
    friend ostream &operator<<(ostream &os, const modint &m) {
        return os << m.n;
    }
};

using mi = modint;

struct Comb {
    vector<mi> fact, inv_fact;
    Comb(int n) : fact(n + 1), inv_fact(n + 1) {
        fact[0] = 1;
        for (int i = 1; i <= n; i++) {
            fact[i] = fact[i - 1] * i;
        }

        inv_fact[n] = fact[n].inv();
        for (int i = n - 1; i >= 0; i--) {
            inv_fact[i] = inv_fact[i + 1] * (i + 1);
        }
    }
    mi C(int n, int k) {
        if (n < 0 || k < 0 || n < k) return 0;
        return fact[n] * inv_fact[k] * inv_fact[n - k];
    }

    mi P(int n, int k) {
        if (n < 0 || k < 0 || n < k) return 0;
        return fact[n] * inv_fact[n - k];
    }

    mi H(int n, int k) {
        if (n < 0 || k < 0) return 0;
        return C(n + k - 1, k);
    }

    mi multinomial(int n, const vector<int> &ks) {
        mi res = fact[n];
        for (int k : ks) {
            if (k < 0 || k > n) return 0;
            res *= inv_fact[k];
        }
        return res;
    }

    mi catalan(int n) { return C(2 * n, n) * mi(n + 1).inv(); }
};

void solve() {
    int n, m;
    cin >> n >> m;
    vector<int> c(n);
    REP(i, n) cin >> c[i];
    Comb comb(max(n + 1, m + 1));
    mi base = 0;
    vector<mi> vio(m + 1, 0);
    auto update = [&](int s, int k, mi c, int sign) {
        base += sign * modint(m).pow(s) * c;
        for (int ban = 1; ban <= k; ban++) {
            mi t = comb.C(k, ban) * modint(m - ban).pow(s);
            vio[ban] += sign * t * c;
        }
    };
    mi ans = 0;
    vector<int> left_0(n, 0);
    REP(i, n) left_0[i] = (i == 0 ? 0 : left_0[i - 1]) + (c[i] == 0);
    int all_0 = count(ALL(c), 0);
    REP(i, n) {
        if (c[i] == 0) {
            base *= m;
            REP(j, m + 1) vio[j] *= m - j;
            update(1, m, modint(m).pow(left_0[i] - 1), 1);
        } else {
            int num_0 = 0;
            set<int> st;
            for (int j = i - 1; j >= 0; j--) {
                if (c[j] == c[i]) break;
                if (c[j] == 0) num_0++;
                else st.insert(c[j]);
                update(num_0, m - st.size(),
                       modint(m).pow(left_0[j] - (c[j] == 0)), -1);
                update(num_0, m - st.size() - 1,
                       modint(m).pow(left_0[j] - (c[j] == 0)), 1);
            }
            update(0, m - 1, modint(m).pow(left_0[i]), 1);
        }
        mi tmp = base;
        for (int ban = 1; ban <= m; ban++) {
            if (ban & 1) tmp -= vio[ban];
            else tmp += vio[ban];
        }
        ans += tmp * modint(m).pow(all_0 - left_0[i]);
    }
    cout << ans << endl;
}

int main() {
    FASTIO();
    int t = 1;
    while (t--) {
        solve();
    }
}
0