切片与 target_size

把内部腿的所有取值拆成多个精确子任务,并准确解释单片中间张量上限。

本页目录

切片是精确分解

选择若干内部腿后,ArcTN 枚举这些腿的全部取值组合。每个 assignment 固定相应输入轴,沿同一条 SSA path 完成一次收缩,最后把全部结果逐元素相加。切片在数学上是精确分解:它不做低秩截断,也不遗漏 assignment。实际使用浮点数时,切片求和与未切片收缩可能采用不同的加法顺序,因此可能存在正常的舍入差异,不承诺逐位相同。

图中用一条维数为 5 的内部腿说明切片恒等式。切片降低每个子任务的中间张量规模,但通常增加总工作量和调度开销;被切的具体腿由网络与路径决定。

若切片腿维度为 d1, ..., dk,切片数是它们的乘积。SliceResult.log2_n_slices 保存该数量的 base-2 logarithm;它不要求所有腿都是二进制维度。

target_size 的严格含义

公开 target_size 是严格正整数,单位是元素数。它要求每个 slice 中路径的各个二元收缩结果张量都不超过该元素数;单张量网络则检查一元处理后的最终输出,并使用整数算术复核。它不限制一元预处理产生的全部临时张量,也不是进程 RSS 或 GPU 显存上限。

target_size 约束 target_size 不约束
每个 slice 的最大二元收缩结果元素数;单张量网络则检查最终输出 进程 RSS 或 GPU 显存
路径与 sliced legs 配对后的二元收缩结果 所有 live tensors 的元素数之和
逻辑元素数 dtype 换算后的字节数
单个 slice 多个并发 slice 的合计驻留内存
结果张量 |C| 单步 |A| + |B| + |C| 或后端 workspace

Warning

三个不同的量

log2_max_contraction_size 报告最大单步逻辑输入加输出元素数的以 2 为底对数,其中原始叶张量按一元求和、迹或对角处理后的逻辑大小计数;log2_peak_size 报告路径模型中存活元素数峰值的以 2 为底对数。它们都不是 target_size 的另一种写法。

单张量网络的 SSA path 为空,但仍会发生最终的迹、求和或输出轴整理。实现会检查 unary-processed output root,不能因为路径没有二元步骤就把过小的 target_size 判为可行。

固定路径的切片腿选择

find_slices_to_size 保持路径不变,把已经选择的腿维度临时设为 1 并重放路径。若仍有中间结果超过目标,下一条腿按以下顺序选择:覆盖更多超目标中间结果者优先;其次选择维度更大的腿;最后用较小腿 ID 作确定性 tie-break。

三个超过目标大小的中间张量都含有 x,y、z 各出现两次,因此先选择 x。随后将 x 的有效维度设为 1,沿原路径重新计算中间张量大小。
  1. 切片腿必须在网络和 size_dict 中存在,且维度为正。
  2. 输出腿不能切片,因为它必须保留在结果中。
  3. 切片列表中不能重复同一条腿。
  4. 同一输入张量中重复出现的腿不能切片;这种腿具有 trace 或 diagonal 语义。
  5. 没有合法腿可以继续降低结果规模时,目标不可达。
rust
use arctn::{find_slices_to_size, slice_result_fits_target_size};

let target_size = 1usize << 24;
let slices = find_slices_to_size(&net, &path, target_size)
    .ok_or("target_size is unreachable for this fixed path")?;
assert!(slice_result_fits_target_size(
    &net, &path, &slices, target_size
)?);

选择规则的保证边界

这是固定路径上的贪心腿选择规则。它不保证用最少的 sliced legs,也不保证切片后的总 FLOPs 最低;最终只保证返回方案通过目标约束复核,或明确报告不可达。

读取 SliceResult

字段 含义
legs 实际选择的内部腿 ID
log2_n_slices 全部 assignment 数量的 log2
per_slice 把 sliced legs 维度设为 1 后的 PathStats
log10_flops_total per-slice FLOPs 乘切片数后的总 FLOPs 的以 10 为底对数

最大中间结果是 per-slice 指标,不乘切片数;总 FLOPs 和总 read/write 工作才要包含切片数乘数。评估一个方案时,至少同时记录路径、sliced legs、切片数、per-slice 最大结果、总 FLOPs 和实际执行时间。

路径与切片计划配对

路径和 SliceResult 必须成对保存。允许改变路径的切片方法会同时返回对应的最终路径;不能把一条路径产生的 sliced legs 无检查地配到另一条路径上。