E - 巡回パスワード / Cyclic Password 解説 by admin
Gemini 3.0 Flash (Thinking)概要
\(N\) 台の端末を \(K\) 回移動しながら、各端末で記録した数字を並べてできる大きな整数を \(M\) で割った余りを求める問題です。操作回数 \(K\) が最大 \(10^{18}\) と非常に大きいため、愚直なシミュレーションではなく、ダブリング(Binary Lifting)を用いた高速化が必要となります。
考察
1. 値の計算方法の整理
\(K\) 回の操作で記録された数字を \(d_1, d_2, \ldots, d_K\) としたとき、得られる値 \(V\) は以下の式で表されます。 $\(V = d_1 \cdot 10^{K-1} + d_2 \cdot 10^{K-2} + \dots + d_K \cdot 10^0\)$
この計算式は、新しい数字 \(d_{new}\) を末尾に付け加える操作を考えると分かりやすくなります。 現在の値を \(V_{old}\)、その長さ(桁数)を \(L\) とすると、さらに \(k\) 桁の数値 \(V_{add}\) を後ろに連結した後の新しい値 \(V_{new}\) は以下のようになります。 $\(V_{new} = V_{old} \cdot 10^k + V_{add} \pmod M\)$
2. なぜ愚直なシミュレーションではダメか
各クエリに対して \(K\) 回の移動をそのままシミュレーションすると、1クエリあたり \(O(K)\) の時間がかかります。\(K \leq 10^{18}\) であるため、これでは制限時間内に終わりません。
3. ダブリングによる高速化
「\(2^p\) 回移動した後の場所」と「その時に得られる値」をあらかじめ計算しておくことで、任意の \(K\) 回の移動を \(O(\log K)\) で処理できるようになります。
next_node[p][i]:端末 \(i\) から出発して \(2^p\) 回移動した後の端末番号value[p][i]:端末 \(i\) から出発して \(2^p\) 回移動したときに得られる値(\(\pmod M\))
これらが分かれば、p+1 の状態は p の状態を2回分組み合わせることで計算できます。
- 移動先: next_node[p+1][i] = next_node[p][next_node[p][i]]
- 値: value[p+1][i] = (value[p][i] * 10^{2^p} + value[p][next_node[p][i]]) % M
アルゴリズム
前処理(ダブリングテーブルの構築)
- \(p=0\)(\(2^0=1\) 回移動)のとき、
next_node[0][i] = P_i、value[0][i] = D_i \pmod Mです。 - \(p=1, 2, \dots, \log_2(\max K)\) について、上記の遷移式を用いてテーブルを埋めます。
- 同時に、\(10^{2^p} \pmod M\) も計算しておきます。
- \(p=0\)(\(2^0=1\) 回移動)のとき、
クエリ処理
- 各クエリ \((S_j, K_j)\) について、現在の値を \(v = 0\)、現在の端末を \(u = S_j\) とします。
- \(K_j\) を 2 進数展開し、ビットが立っている桁 \(p\) について以下の更新を行います。
- \(v = (v \cdot 10^{2^p} + \text{value}[p][u]) \pmod M\)
- \(u = \text{next\_node}[p][u]\)
- 最終的な \(v\) が答えとなります。
計算量
- 時間計算量: \(O((N + Q) \log (\max K))\)
- ダブリングテーブルの構築に \(O(N \log K)\)、クエリ応答に \(O(Q \log K)\) かかります。
- 空間計算量: \(O(N \log (\max K))\)
- ダブリングテーブルを保持するためのメモリが必要です。
実装のポイント
メモリ節約: Python の場合、
listのリストでテーブルを持つとメモリを大量に消費するため、array.arrayを使用したり、必要最低限のビット数でループを回したりする工夫が有効です。高速化: \(10^{2^p} \pmod M\) は繰り返し二乗法の要領で事前に計算しておくと、クエリごとの計算を減らせます。
0-indexed: 端末番号が \(1\) から \(N\) で与えられるため、実装上は \(0\) から \(N-1\) に変換しておくと配列のアクセスがスムーズです。また、\(K=0\) のケース(答えは \(0\))に注意してください。
ソースコード
import sys
import array
def solve():
# Read all input at once for speed
input_data = sys.stdin.read().split()
if not input_data:
return
it = iter(input_data)
try:
N = int(next(it))
Q = int(next(it))
M = int(next(it))
except StopIteration:
return
# Store digits and next terminal destinations
D = [0] * N
P = [0] * N
for i in range(N):
try:
D[i] = int(next(it))
P[i] = int(next(it)) - 1 # Convert to 0-indexed
except StopIteration:
break
# Store queries and find the maximum number of steps needed
queries = []
max_k = 0
for _ in range(Q):
try:
s = int(next(it)) - 1 # Convert to 0-indexed
k = int(next(it))
queries.append((s, k))
if k > max_k:
max_k = k
except StopIteration:
break
# Special case where all query steps are 0
if max_k == 0:
sys.stdout.write('\n'.join(['0'] * Q) + '\n')
return
# Calculate the number of bits needed for the maximum k
num_bits = max_k.bit_length()
# next_node[bit][i] stores the terminal reached after 2^bit steps from terminal i
# value[bit][i] stores the value obtained after 2^bit steps from terminal i
# We use array.array for memory efficiency
next_node = [array.array('I', P)]
value = [array.array('I', [d % M for d in D])]
# pow10[bit] stores 10^(2^bit) % M
pow10 = array.array('Q', [0] * num_bits)
pow10[0] = 10 % M
# Binary lifting precomputation: O(N log K)
for bit in range(num_bits - 1):
prev_next = next_node[bit]
prev_value = value[bit]
p10 = pow10[bit]
pow10[bit + 1] = (p10 * p10) % M
# next_node[bit+1][i] = next_node[bit][next_node[bit][i]]
# value[bit+1][i] = (value[bit][i] * 10^(2^bit) + value[bit][next_node[bit][i]]) % M
# Using list comprehensions for faster execution in Python
next_node.append(array.array('I', [prev_next[mid] for mid in prev_next]))
value.append(array.array('I', [(int(v) * p10 + prev_value[mid]) % M
for v, mid in zip(prev_value, prev_next)]))
results = []
# Process each query: O(Q log K)
for s, k in queries:
v = 0
u = s
bit = 0
# Decompose k into powers of 2 and jump through the precomputed tables
while k:
if k & 1:
# Update current value: shift existing value by 2^bit places and add new part
v = (v * pow10[bit] + value[bit][u]) % M
# Jump to the terminal reached after 2^bit steps
u = next_node[bit][u]
k >>= 1
bit += 1
results.append(v)
# Output all results at once
sys.stdout.write('\n'.join(map(str, results)) + '\n')
if __name__ == '__main__':
solve()
この解説は gemini-3-flash-thinking によって生成されました。
投稿日時:
最終更新: