G - Caeser Syllables Editorial
by
sheyasutaka
本稿では,多項式 \(P\) における \(x^i\) の係数を \([x^i] P\) と表記し,この表記を多変数多項式にも拡張します.
言い換え
母音の集合を \(\mathbf{V}\) と表記します.
記号列 \(A\) において,音節の先頭の要素 \(A_v\) は,以下のいずれかを満たします.
- \(v = 1\) であり,\(A_v \in \mathbf{V}\)
- \(v \geq 2\) であり,\(A_{v-1} \notin \mathbf{V}\) かつ \(A_v \in \mathbf{V}\)
上のいずれかを満たす \(v\) の個数が \(A\) の音節数に等しいことから,そのような \(v\) の個数を求める問題に言い換えられます.
\(v = 1\) のときに満たすかは容易に求まるので,\(v \geq 2\) における個数を求める方法を考えます.
\((A_{v-1} + k) \bmod K \notin \mathbf{V}\) かつ \((A_v + k) \bmod K \in \mathbf{V}\) を満たす \(v \geq 2\) の個数
愚直に求めようとすると \(\Theta(K \cdot \min(K^2, N))\) 時間かかるので,より高速に求める必要があります.
\(x,y\) についての \(2\) 変数多項式 \(A,S,T\) を以下のように定めます.
\[A := \sum_{2 \leq v \leq N} x^{K - A_{v-1}} y^{K - A_{v}}\]
\[S := \sum_{\substack{i \notin \mathbf{V} \\ 0 \leq i \leq K-1}} x^i \]
\[T := \sum_{\substack{j \in \mathbf{V} \\ 0 \leq j \leq K-1}} y^j \]
このとき,\((A_{v-1} + k) \bmod K \notin \mathbf{V}\) かつ \((A_v + k) \bmod K \in \mathbf{V}\) を満たす \(v\) の個数は \(\displaystyle \sum_{\substack{i,j \in \{k,\ k+K\}}} [x^i y^j] AST\) に等しくなります.したがって,\(AST\) の特定の次数 \(O(K)\) 通りの係数を求めることができれば十分です.
\(A,S,T\) ならびに \(AST\) の各係数は \(N\) 以下の非負整数になります.したがって,十分大きな素数 \(p\) をとって \(\mathbb{Z}/p\mathbb{Z}\) 上の数論変換 (NTT) で \(AST\) の各係数を求めることができれば,それが一意に \(AST\) の真の係数を示します.
\(AST\) の求め方として,複数の方法が考えられます.
方法 1: \(S,T\) の形を利用し,\(AS\),\(AST\) を順に求める
\(S\) は \(x\) のみを,\(T\) は \(y\) のみを変数にもつ \(1\) 変数多項式として扱えます.このとき,以下の \(2\) ステップによって,\(A \times S\) および \(A \times S \times T\) の各係数を順に求めることができます.
- \(x\) の \(1\) 変数多項式 \(A^{(j)}\) を \(A^{(j)}(x) := [y^j]A(x,y)\) として定める.このとき,\([y^j](A \times S) = A^{(j)} \times S\) である.
- \(y\) の \(1\) 変数多項式 \((A \times S)^{(i)}\) を \((A \times S)^{(i)}(y) := [x^i](A \times S)(x,y)\) として定める.このとき,\([x^i](A \times S \times T) = (A \times S)^{(i)} \times T\) である.
\(K\) 次の \(1\) 変数多項式どうしの積を \(O(K)\) 回求めればよく,これは一般的な (\(p=998244353\) の) 数論変換で実現可能です.時間計算量は \(O(N + K^2 \log K)\) であり,十分高速です.
方法 2: \(AST\) を直接求める
\(y\) を \(x^{2K}\) に置き換えることで,\(A,S,T\) を \(x\) の \(1\) 変数多項式として扱って積 \(AST\) を求めることができます.
\(AST\) は \(4K^2 + 2K\) 次の多項式になり,本問題の制約においてこの次数の最大値は \(21164600\ ({} \simeq 2^{24.34})\) です.よって,数論変換に使う素数 \(p\) は \(2^{25}\ |\ (p-1)\) を満たす必要があることに注意してください(\(998244353\) は満たしません).
時間計算量は \(O(N + 2 K^2 \log K)\) であり,十分高速です.
実装例 (C++)
#include <iostream>
using std::cin;
using std::cout;
using std::cerr;
using std::endl;
#include <vector>
using std::vector;
using std::pair;
#include <map>
using std::map;
using std::max;
using std::min;
#ifdef DEBUG
const int debug = 1;
#else
const int debug = 0;
#endif
using ll = int64_t;
using P = pair<ll, ll>;
const ll FOD = 998244353;
#include <atcoder/modint>
using mint = atcoder::modint998244353;
#include <atcoder/convolution>
using atcoder::convolution;
ll n, k, m;
uint64_t seed;
ll a[7'000'456], v[3'123];
inline void output (const ll x) {
cout << x << "\n";
}
void solve() {
vector<ll> rarr[k]; // rarr[i] := sum of x^{K-A_{i-1}} where K-A_{i-0} == i
for (ll i = 0; i < k; i++) {
rarr[i].assign(k, 0);
}
for (ll i = 1; i < n; i++) {
rarr[(k - a[i-0]) % k][(k - a[i-1]) % k] += 1;
}
vector<ll> s(k, 0), t(k, 0); // s[i] := (i notin VOWEL), t[i] := (i in VOWEL)
for (ll i = 0; i < k; i++) {
if (v[i] == 1) {
t[i] += 1;
} else {
s[i] += 1;
}
}
vector<ll> rsarr[k]; // rsarr[i] := (rarr[i] * s), x^K loops back to x^0
for (ll i = 0; i < k; i++) {
vector<ll> p = convolution(rarr[i], s);
rsarr[i].assign(k, 0);
for (ll v = 0; v < (ll)p.size(); v++) {
rsarr[i][v % k] += p[v];
}
}
vector<ll> rsbrr[k]; // rsbrr[x][y] is the transpose of rsarr[y][x]
for (ll i = 0; i < k; i++) {
rsbrr[i].resize(k);
for (ll j = 0; j < k; j++) {
rsbrr[i][j] = rsarr[j][i];
}
}
vector<ll> rst[k]; // rst[j] := (rsbrr[j] * t), y^K loops back to y^0
for (ll i = 0; i < k; i++) {
vector<ll> q = convolution(rsbrr[i], t);
rst[i].assign(k, 0);
for (ll v = 0; v < (ll)q.size(); v++) {
rst[i][v % k] += q[v];
}
}
// now rst[i][j] is [x^i y^j] (sum of x^{-A_{i-1}} y^{-A_{i-0}}) * (sum of x^NONVOWEL) * (sum of y^VOWEL)
// rst[k][k] + (A_0 + k is a vowel ? 1 : 0) is the answer for k
for (ll i = 0; i < k; i++) {
ll ans = rst[i][i];
if (v[(a[0] + i) % k]) ans += 1;
output(ans);
}
return;
}
int main (void) {
std::cin.tie(nullptr);
std::ios_base::sync_with_stdio(false);
cin >> n >> k >> seed >> m;
uint64_t state = seed;
for (ll i = 0; i < n; i++) {
if (i < m) {
cin >> a[i];
} else {
uint64_t x = (((state >> 18) ^ state) >> 27) & ((1ULL << 32) - 1);
uint64_t r = state >> 59;
uint64_t y = ((x >> r) + (x << (32-r))) & ((1ULL << 32) - 1);
a[i] = (y % (uint64_t)k) + 0;
state = (state * 6364136223846793005ULL + 2026081520260815ULL);
}
a[i] -= 0;
}
for (ll i = 0; i < k; i++) {
cin >> v[i];
}
solve();
return 0;
}
posted:
last update:
