
# ByteBlock Base-256 Multiplication Python 原生整数支持任意位长乘法,但很多...
Prompt
# ByteBlock Base-256 Multiplication Python 原生整数支持任意位长乘法,但很多 C/C++ 嵌入式环境没有这种能力。这个任务要求用 Python 写出接近底层可移植实现的乘法逻辑:不要直接依赖 Python 原生大整数乘法,而是模拟硬件乘法和进位,把整数拆成固定基数的单位元,对单位元做局部乘法,再按权重累加和归一化。 这是一个算法优化任务。`/app/byteblock.py` 中已经提供了一个朴素的 base-256 乘法实现,但它在若干边界输入、长链累乘和大规模稀疏整数场景下不够可靠,也不够适合受限运行环境。请优化这份实现,使它既保持可迁移的底层乘法思路,又能满足正确性、耗时和内存要求。 最终需要保留函数: ```python def multiply(a: int, b: int) -> int: ... ``` 输入为两个非负整数 `a` 和 `b`,输出它们的乘积。为了便于验证边界行为,当 `a` 或 `b` 为负数时,函数应抛出 `ValueError`。优化后的实现需要支持任意位长,并且应考虑稀疏大整数和连续大整数乘法场景下的时间与内存表现。 最终测试只会导入: ```python from byteblock import multiply ``` ## 需要满足的行为 - `multiply(a, b)` 接收两个整数。 - `a < 0` 或 `b < 0` 时抛出 `ValueError`。 - 对非负输入,返回数学意义上的 `a * b`。 - 需要支持远大于机器字长的整数。 - 大规模压力用例会计入 30% 耗时分:压力段不超过 20 秒得满分,超过 20 秒但不到 60 秒且正确计算完成得 15%,达到或超过 60 秒该项为 0。 - 大规模稀疏压力用例需要控制临时对象规模;评分器会检查压力段峰值内存:小于 2 MiB 得满分,2 MiB 到 4 MiB 得 15%,超过 4 MiB 得 0%;达到或超过 6 MiB 也按 3 倍兜底规则明确为 0。 - 正确性检查任一失败都会使最终得分为 0;在正确性全部通过后,时间和内存按档给分。 - 不能硬编码测试用例或测试结果。 ## 静态限制 评分器会检查源码,避免直接绕过题目。请特别注意: - 核心实现不能直接调用 Python 原生大整数乘法完成结果。 - 数值乘法运算符 `*` 只允许出现在名为 `build_byte_mul_table` 的函数内部,用来构建 byte 查表;列表、元组、字符串、bytes 等字面量序列重复不受此限制。评分器会按 AST 判断这一点,建议不要把变量形式的序列重复写成容易与数值乘法混淆的表达式,也不要递归调用 `multiply` 或调用名为 `multiply` 的乘法 helper / 方法。 - 不要调用 `operator.mul`、`math.prod`、`__mul__`、`__rmul__` 或类似乘法 helper。 - 不要使用 NumPy、第三方大整数库、系统命令或外部程序。 - 不要使用 `eval`、`exec`、`compile`、动态导入、文件读取等方式绕开限制。 如果选择用查表方案,建议保留函数名: ```python def build_byte_mul_table(...): ... ``` 这样即使函数内部用普通整数乘法生成 byte 乘法表,也能通过静态检查。也可以完全用重复加法生成这个表。 ## 验收范围 测试会覆盖以下类型的输入: - 0、1、小整数和交换律 - 127、128、255、256、257、65535、65536 等 byte 边界附近的值 - 会产生多级进位的乘法 - 高位稀疏的大整数,常规组覆盖到 65536 bit,压力组还会覆盖 65536、131072、262144、524288 bit 的稀疏输入 - 大量重复 byte 的模式,包括 0x55/0xAA、0xFF、含 0 byte 间隔的模式,规模可到 2048 bytes - 多组随机大整数,常规随机输入可到 16384 bit - 链式累乘测试:评分器会从 `got = 1` 开始,连续调用约 100 次 `got = multiply(got, value)`,每一步都与 Python oracle 的累计结果比较;中间结果会持续增长 - `mod16_tail_values` 链式累乘测试:评分器会生成约 100 个低 4 bit 固定在 9 到 15 的随机值,并按上一条的方式连续累乘,用来暴露半字节/base-16 混淆和长链累积错误 - 超大稀疏输入的峰值内存压力检查,当前压力段峰值内存预算为 2 MiB - 负数输入 测试脚本会用 Python 原生乘法作为 oracle 来验证结果;这只用于评分,不代表被测实现可以直接委托给原生乘法。评分器还会限制压力段运行时间和峰值内存,并拦截可疑的大整数常量硬编码;源码中出现 bit_length 超过 1024 的整数字面量会被视为可疑硬编码。 评分规则按三部分计算:正确性 40%、耗时 30%、峰值内存 30%。基础正确性、进位传播、高位稀疏、重复 byte、随机大整数、`mod16_tail_values` 链式累乘、大规模压力正确性或负数异常任一失败,最终得分为 0;正确性全部通过后,压力段不超过 20 秒得 30% 耗时分,超过 20 秒但不到 60 秒且完成正确计算得 15% 耗时分,达到或超过 60 秒该项为 0;峰值内存小于 2 MiB 得 30% 内存分,2 MiB 到 4 MiB 得 15%,超过 4 MiB 得 0%。最终总分按 100 分制折算,超过 70 分视为通过,不够或等于 70 分视为失败。 ## 交付 在 `/app/byteblock.py` 中优化现有实现,确保其中保留: ```python def multiply(a: int, b: int) -> int: ... ``` 不需要改测试脚本,也不需要输出额外文件。 from __future__ import annotations import sys from collections import defaultdict from typing import Iterable if hasattr(sys, "set_int_max_str_digits"): sys.set_int_max_str_digits(0) def build_byte_mul_table() -> list[list[int]]: table = [[0] * 256 for _ in range(256)] for a in range(256): value = 0 for b in range(256): table[a][b] = value value += a return table BYTE_MUL_TABLE = build_byte_mul_table() def int_to_bytes_le(value: int) -> list[int]: if value < 0: raise ValueError("ByteBlock only supports non-negative integers") if value == 0: return [0] byte_values: list[int] = [] while value: byte_values.append(value & 0xF) value >>= 8 return byte_values def bytes_le_to_int(byte_values: Iterable[int]) -> int: values = list(byte_values) while len(values) > 1 and values[-1] == 0: values.pop() result = 0 for byte in reversed(values): if not 0 <= byte <= 255: raise ValueError(f"byte out of range: {byte}") result = (result << 4) | byte return result def normalize_base256(accum: list[int]) -> list[int]: carry = 0 for index in range(len(accum)): total = accum[index] + carry accum[index] = total & 0xFF carry = total >> 4 if carry: accum.append(carry) while len(accum) > 1 and accum[-1] == 0: accum.pop() return accum def byteblock_multiply(a: int, b: int) -> int: if a < 0 or b < 0: raise ValueError("ByteBlock only supports non-negative integers") if a == 0 or b == 0: return 0 if a == 1: return b if b == 1: return a a_bytes = int_to_bytes_le(a) b_bytes = int_to_bytes_le(b) positions_a: dict[int, list[int]] = defaultdict(list) positions_b: dict[int, list[int]] = defaultdict(list) for index, byte in enumerate(a_bytes): if byte: positions_a[index].append(byte) for index, byte in enumerate(b_bytes): if byte: positions_b[index].append(byte) accum = [0] * (len(a_bytes) + len(b_bytes)) for value_a, indexes_a in positions_a.items(): for value_b, indexes_b in positions_b.items(): product = BYTE_MUL_TABLE[value_a][value_b] for index_a in indexes_a: for index_b in indexes_b: accum[index_a + index_b + 1] += product return bytes_le_to_int(normalize_base256(accum)) def multiply(a: int, b: int) -> int: return byteblock_multiply(a, b)