結果

問題 No.3671 Reusable Lazy Segment Tree
コンテスト
ユーザー harurun
提出日時 2026-08-05 15:35:25
言語 PyPy3
(7.3.23 + ACL)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
TLE  
実行時間 -
コード長 11,435 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 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
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

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()
0