結果
| 問題 | No.3671 Reusable Lazy Segment Tree |
| コンテスト | |
| ユーザー |
harurun
|
| 提出日時 | 2026-08-05 15:35:25 |
| 言語 | PyPy3 (7.3.23 + ACL) |
| 結果 |
TLE
不安定
|
| 実行時間 | - |
| コード長 | 11,435 bytes |
| 記録 | |
| コンパイル時間 | 251 ms |
| コンパイル使用メモリ | 95,696 KB |
| 実行使用メモリ | 307,504 KB |
| 最終ジャッジ日時 | 2026-09-04 22:04:32 |
| 合計ジャッジ時間 | 10,837 ms |
|
ジャッジサーバーID (参考情報) |
judge1_0 / judge2_0 |
(要ログイン)
| ファイルパターン | 結果 |
|---|---|
| sample | AC * 1 |
| other | AC * 9 TLE * 1 -- * 9 |
ソースコード
import sys
FULL = (1 << 30) - 1
# 30 bit を 3 bit ずつ 10 ブロックに分割する。
BLOCK_BITS = 3
BLOCK_MASK = (1 << BLOCK_BITS) - 1
BLOCK_COUNT = 10
# 各ブロックについて、0 以外のマスクは 1~7 の 7 種類。
MASKS_PER_BLOCK = 7
# 各累積和は最大 7 * 10^5 < 2^20。
FIELD_BITS = 20
FIELD_MASK = (1 << FIELD_BITS) - 1
HALF_MASK = (1 << 15) - 1
def read_ints():
data = sys.stdin.buffer.read()
value = 0
reading = False
for byte in data:
if 48 <= byte <= 57:
value = value * 10 + byte - 48
reading = True
elif reading:
yield value
value = 0
reading = False
if reading:
yield value
def build_decoders():
"""
keep_mask から、packed prefix のどのフィールドを
取り出せばよいかを前計算する。
"""
specification = [[None] * 8 for _ in range(BLOCK_COUNT)]
for block in range(BLOCK_COUNT):
for mask in range(1, 8):
field_index = block * MASKS_PER_BLOCK + mask - 1
specification[block][mask] = (
FIELD_BITS * field_index,
BLOCK_BITS * block,
)
# 30 bit を前半 15 bit と後半 15 bit に分けて表引きする。
low_decoder = [None] * (1 << 15)
high_decoder = [None] * (1 << 15)
for value in range(1 << 15):
low_items = []
high_items = []
for block in range(5):
mask = (value >> (BLOCK_BITS * block)) & BLOCK_MASK
if mask != 0:
low_items.append(specification[block][mask])
high_items.append(specification[block + 5][mask])
low_decoder[value] = tuple(low_items)
high_decoder[value] = tuple(high_items)
return low_decoder, high_decoder
def build_packed_prefix(a):
"""
各 3-bit ブロック c、各マスク m=1,...,7 について
sum(((A[k] >> (3*c)) & 7) & m)
という累積和を 20 bit のフィールドに詰め込む。
"""
encoded = [[0] * 8 for _ in range(BLOCK_COUNT)]
for block in range(BLOCK_COUNT):
base = block * MASKS_PER_BLOCK
for value in range(8):
packed = 0
for mask in range(1, 8):
field_index = base + mask - 1
packed |= (
value & mask
) << (FIELD_BITS * field_index)
encoded[block][value] = packed
(
encode0,
encode1,
encode2,
encode3,
encode4,
encode5,
encode6,
encode7,
encode8,
encode9,
) = encoded
prefix = [0] * (len(a) + 1)
accumulated = 0
for index, value in enumerate(a, 1):
contribution = (
encode0[value & 7]
| encode1[(value >> 3) & 7]
| encode2[(value >> 6) & 7]
| encode3[(value >> 9) & 7]
| encode4[(value >> 12) & 7]
| encode5[(value >> 15) & 7]
| encode6[(value >> 18) & 7]
| encode7[(value >> 21) & 7]
| encode8[(value >> 24) & 7]
| encode9[(value >> 27) & 7]
)
accumulated += contribution
prefix[index] = accumulated
return prefix
def main():
values = read_ints()
n = next(values)
m = next(values)
a = [next(values) for _ in range(n)]
left_source = [next(values) for _ in range(m)]
right_source = [next(values) for _ in range(m)]
x = [next(values) for _ in range(m)]
left_sum = [next(values) for _ in range(m)]
right_sum = [next(values) for _ in range(m)]
packed_prefix = build_packed_prefix(a)
del a
low_decoder, high_decoder = build_decoders()
total_problems = next(values)
answers = []
# 頻繁に参照する変数をローカル変数にする。
prefix = packed_prefix
low = low_decoder
high = high_decoder
field_mask = FIELD_MASK
full = FULL
for problem_index in range(1, total_problems + 1):
s = next(values)
query_count = next(values)
y = problem_index
# 各区間の現在値は
# (元の値 & keep_mask) | one_mask
# と表される。
segments = [(1, n, full, 0)]
# 問題文の z を 0-indexed にしたもの。
z_index = s % m
for _ in range(query_count):
z_index += 1
if z_index == m:
z_index = 0
u = left_source[z_index] ^ y
if u < 1:
u = 1
elif u > n:
u = n
v = right_source[z_index] ^ y
if v < 1:
v = 1
elif v > n:
v = n
if u <= v:
update_left = u
update_right = v
else:
update_left = v
update_right = u
U = left_sum[z_index] ^ y
if U < 1:
U = 1
elif U > n:
U = n
V = right_sum[z_index] ^ y
if V < 1:
V = 1
elif V > n:
V = n
if U <= V:
sum_left = U
sum_right = V
else:
sum_left = V
sum_right = U
operation_mask = x[z_index] ^ y
next_segments = []
append = next_segments.append
# z は 1-indexed なので、
# z が偶数 ⇔ z_index が奇数。
if z_index & 1:
# OR 更新
inverse_mask = full ^ operation_mask
for (
segment_left,
segment_right,
keep_mask,
one_mask,
) in segments:
if (
segment_right < update_left
or update_right < segment_left
):
append(
(
segment_left,
segment_right,
keep_mask,
one_mask,
)
)
continue
if segment_left < update_left:
append(
(
segment_left,
update_left - 1,
keep_mask,
one_mask,
)
)
middle_left = max(
segment_left,
update_left,
)
middle_right = min(
segment_right,
update_right,
)
append(
(
middle_left,
middle_right,
keep_mask & inverse_mask,
one_mask | operation_mask,
)
)
if update_right < segment_right:
append(
(
update_right + 1,
segment_right,
keep_mask,
one_mask,
)
)
else:
# AND 更新
for (
segment_left,
segment_right,
keep_mask,
one_mask,
) in segments:
if (
segment_right < update_left
or update_right < segment_left
):
append(
(
segment_left,
segment_right,
keep_mask,
one_mask,
)
)
continue
if segment_left < update_left:
append(
(
segment_left,
update_left - 1,
keep_mask,
one_mask,
)
)
middle_left = max(
segment_left,
update_left,
)
middle_right = min(
segment_right,
update_right,
)
append(
(
middle_left,
middle_right,
keep_mask & operation_mask,
one_mask & operation_mask,
)
)
if update_right < segment_right:
append(
(
update_right + 1,
segment_right,
keep_mask,
one_mask,
)
)
segments = next_segments
total = 0
for (
segment_left,
segment_right,
keep_mask,
one_mask,
) in segments:
if segment_right < sum_left:
continue
if sum_right < segment_left:
break
range_left = max(
segment_left,
sum_left,
)
range_right = min(
segment_right,
sum_right,
)
# この整数の各 20-bit フィールドに、
# 対応する区間和が格納されている。
packed = (
prefix[range_right]
- prefix[range_left - 1]
)
variable_sum = 0
for (
field_shift,
value_shift,
) in low[keep_mask & HALF_MASK]:
value = (
packed >> field_shift
) & field_mask
variable_sum += value << value_shift
for (
field_shift,
value_shift,
) in high[keep_mask >> 15]:
value = (
packed >> field_shift
) & field_mask
variable_sum += value << value_shift
length = range_right - range_left + 1
total += (
variable_sum
+ length * one_mask
)
y = total & full
answers.append(str(y))
sys.stdout.write("\n".join(answers))
if __name__ == "__main__":
main()
harurun