結果

問題 No.3676 Cuboid Alignment
コンテスト
ユーザー 👑 みうね
提出日時 2026-08-10 17:43:56
言語 PyPy3
(7.3.23)
コンパイル:
pypy3 -mpy_compile _filename_
実行:
pypy3 _filename_
結果
AC  
実行時間 1,569 ms / 2,000 ms
+ 853µs
コード長 9,001 bytes
記録
記録タグの例:
初AC ショートコード 純ショートコード 純主流ショートコード 最速実行時間
コンパイル時間 626 ms
コンパイル使用メモリ 96,084 KB
実行使用メモリ 166,828 KB
最終ジャッジ日時 2026-09-04 22:18:51
合計ジャッジ時間 30,851 ms
ジャッジサーバーID
(参考情報)
judge3_0 / judge1_0
このコードへのチャレンジ
(要ログイン)
ファイルパターン 結果
sample AC * 4
other AC * 42
権限があれば一括ダウンロードができます

ソースコード

diff #
raw source code

#!/usr/bin/env pypy3

import sys


MOD = 998244353
PRIMITIVE_ROOT = 3
BLACK = ord('B')
WHITE = ord('W')


def make_ntt_constants():
    max_base = 23
    root = [1] * (max_base + 1)
    inverse_root = [1] * (max_base + 1)
    for level in range(1, max_base + 1):
        root[level] = pow(PRIMITIVE_ROOT, (MOD - 1) >> level, MOD)
        inverse_root[level] = pow(root[level], MOD - 2, MOD)

    rate = [1] * max_base
    inverse_rate = [1] * max_base
    product = inverse_product = 1
    for level in range(max_base - 1):
        rate[level] = root[level + 2] * product % MOD
        inverse_rate[level] = inverse_root[level + 2] * inverse_product % MOD
        product = product * inverse_root[level + 2] % MOD
        inverse_product = inverse_product * root[level + 2] % MOD

    rate4 = [1] * max_base
    inverse_rate4 = [1] * max_base
    product = inverse_product = 1
    for level in range(max_base - 2):
        rate4[level] = root[level + 3] * product % MOD
        inverse_rate4[level] = inverse_root[level + 3] * inverse_product % MOD
        product = product * inverse_root[level + 3] % MOD
        inverse_product = inverse_product * root[level + 3] % MOD
    return root, inverse_root, rate, inverse_rate, rate4, inverse_rate4


ROOT, INVERSE_ROOT, RATE, INVERSE_RATE, RATE4, INVERSE_RATE4 = make_ntt_constants()


def trailing_zeros(value):
    return (value & -value).bit_length() - 1


def ntt(values):
    modulus = MOD
    size = len(values)
    height = size.bit_length() - 1
    phase = 0
    while phase < height:
        remaining = height - phase
        if remaining == 1:
            width = 1 << (remaining - 1)
            twiddle = 1
            for block in range(1 << phase):
                offset = block << remaining
                for index in range(offset, offset + width):
                    left = values[index]
                    right = values[index + width] * twiddle % modulus
                    total = left + right
                    values[index] = total if total < modulus else total - modulus
                    difference = left - right
                    values[index + width] = (
                        difference if difference >= 0 else difference + modulus
                    )
                if block + 1 != 1 << phase:
                    twiddle = (
                        twiddle * RATE[trailing_zeros(block + 1)] % modulus
                    )
            phase += 1
            continue

        width = 1 << (remaining - 2)
        twiddle = 1
        imaginary = ROOT[2]
        for block in range(1 << phase):
            twiddle2 = twiddle * twiddle % modulus
            twiddle3 = twiddle2 * twiddle % modulus
            offset = block << remaining
            for index in range(offset, offset + width):
                value0 = values[index]
                value1 = values[index + width] * twiddle % modulus
                value2 = values[index + 2 * width] * twiddle2 % modulus
                value3 = values[index + 3 * width] * twiddle3 % modulus
                value1_minus_value3_i = (value1 - value3) * imaginary % modulus
                values[index] = (value0 + value1 + value2 + value3) % modulus
                values[index + width] = (value0 - value1 + value2 - value3) % modulus
                values[index + 2 * width] = (
                    value0 - value2 + value1_minus_value3_i
                ) % modulus
                values[index + 3 * width] = (
                    value0 - value2 - value1_minus_value3_i
                ) % modulus
            if block + 1 != 1 << phase:
                twiddle = twiddle * RATE4[trailing_zeros(block + 1)] % modulus
        phase += 2


def inverse_ntt(values):
    modulus = MOD
    size = len(values)
    height = size.bit_length() - 1
    phase = height
    while phase:
        if phase == 1:
            width = 1 << (height - phase)
            twiddle = 1
            for block in range(1 << (phase - 1)):
                offset = block << (height - phase + 1)
                for index in range(offset, offset + width):
                    left = values[index]
                    right = values[index + width]
                    total = left + right
                    values[index] = total if total < modulus else total - modulus
                    values[index + width] = (left - right) * twiddle % modulus
                if block + 1 != 1 << (phase - 1):
                    twiddle = (
                        twiddle * INVERSE_RATE[trailing_zeros(block + 1)] % modulus
                    )
            phase -= 1
            continue

        width = 1 << (height - phase)
        twiddle = 1
        inverse_imaginary = INVERSE_ROOT[2]
        for block in range(1 << (phase - 2)):
            twiddle2 = twiddle * twiddle % modulus
            twiddle3 = twiddle2 * twiddle % modulus
            offset = block << (height - phase + 2)
            for index in range(offset, offset + width):
                value0 = values[index]
                value1 = values[index + width]
                value2 = values[index + 2 * width]
                value3 = values[index + 3 * width]
                value2_minus_value3_i = (
                    (value2 - value3) * inverse_imaginary % modulus
                )
                values[index] = (value0 + value1 + value2 + value3) % modulus
                values[index + width] = (
                    (value0 - value1 + value2_minus_value3_i) * twiddle % modulus
                )
                values[index + 2 * width] = (
                    (value0 + value1 - value2 - value3) * twiddle2 % modulus
                )
                values[index + 3 * width] = (
                    (value0 - value1 - value2_minus_value3_i) * twiddle3 % modulus
                )
            if block + 1 != 1 << (phase - 2):
                twiddle = (
                    twiddle * INVERSE_RATE4[trailing_zeros(block + 1)] % modulus
                )
        phase -= 2


def embedded_convolution(
    first_rows,
    second_rows,
    first_color,
    second_color,
    x_size,
    y_size,
    z_size,
    transform_size,
):
    x_radix = 2 * x_size - 1
    y_radix = 2 * y_size - 1
    first_values = [0] * transform_size
    second_values = [0] * transform_size

    row_index = 0
    for z in range(z_size):
        z_base = x_radix * y_radix * z
        reversed_z_base = x_radix * y_radix * (z_size - 1 - z)
        for y in range(y_size):
            first_row = first_rows[row_index]
            second_row = second_rows[row_index]
            first_base = z_base + x_radix * y
            second_base = reversed_z_base + x_radix * (y_size - 1 - y)
            for x in range(x_size):
                if first_row[x] == first_color:
                    first_values[first_base + x] = 1
                if second_row[x] == second_color:
                    second_values[second_base + x_size - 1 - x] = 1
            row_index += 1

    ntt(first_values)
    ntt(second_values)
    modulus = MOD
    for index in range(transform_size):
        first_values[index] = first_values[index] * second_values[index] % modulus
    del second_values
    inverse_ntt(first_values)
    return first_values


def fold_convolution(values, result, x_size, y_size, z_size):
    x_radix = 2 * x_size - 1
    y_radix = 2 * y_size - 1
    inverse_size = pow(len(values), MOD - 2, MOD)
    xy_size = x_size * y_size
    for z in range(2 * z_size - 1):
        source_z = x_radix * y_radix * z
        target_z = xy_size * (z % z_size)
        for y in range(2 * y_size - 1):
            source = source_z + x_radix * y
            target = target_z + x_size * (y % y_size)
            for x in range(x_size):
                result[target + x] += values[source + x] * inverse_size % MOD
            for x in range(x_size - 1):
                result[target + x] += (
                    values[source + x_size + x] * inverse_size % MOD
                )


def solve():
    tokens = sys.stdin.buffer.read().split()
    if not tokens:
        return
    x_size, y_size, z_size = map(int, tokens[:3])
    row_count = y_size * z_size
    first_rows = tokens[3:3 + row_count]
    second_rows = tokens[3 + row_count:]

    embedded_size = (
        (2 * x_size - 1) * (2 * y_size - 1) * (2 * z_size - 1)
    )
    transform_size = 1 << (embedded_size - 1).bit_length()
    result = [0] * (x_size * y_size * z_size)

    convolution = embedded_convolution(
        first_rows,
        second_rows,
        BLACK,
        WHITE,
        x_size,
        y_size,
        z_size,
        transform_size,
    )
    fold_convolution(convolution, result, x_size, y_size, z_size)
    del convolution

    convolution = embedded_convolution(
        first_rows,
        second_rows,
        WHITE,
        BLACK,
        x_size,
        y_size,
        z_size,
        transform_size,
    )
    fold_convolution(convolution, result, x_size, y_size, z_size)
    print(min(result))


if __name__ == '__main__':
    solve()
0