切片与 target_size
把内部腿的所有取值拆成多个精确子任务,并准确解释单片中间张量上限。
本页目录
切片是精确分解
选择若干内部腿后,ArcTN 枚举这些腿的全部取值组合。每个 assignment 固定相应输入轴,沿同一条 SSA path 完成一次收缩,最后把全部结果逐元素相加。切片在数学上是精确分解:它不做低秩截断,也不遗漏 assignment。实际使用浮点数时,切片求和与未切片收缩可能采用不同的加法顺序,因此可能存在正常的舍入差异,不承诺逐位相同。
T[a,d] = Σb,c A[a,b] B[b,c] C[c,d]
c=0T0[a,d]c=1T1[a,d]c=2T2[a,d]c=3T3[a,d]c=4T4[a,d]T[a,d] = T₀ + T₁ + T₂ + T₃ + T₄遍历 c 的全部取值,求和得到原收缩结果。若切片腿维度为 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。
xyzd(x)d(y)d(z)覆盖数降序 → 维度降序 → 腿 ID 升序effective dim(x): d(x) → 1沿同一 SSA path 重新计算中间张量大小- 切片腿必须在网络和 size_dict 中存在,且维度为正。
- 输出腿不能切片,因为它必须保留在结果中。
- 切片列表中不能重复同一条腿。
- 同一输入张量中重复出现的腿不能切片;这种腿具有 trace 或 diagonal 语义。
- 没有合法腿可以继续降低结果规模时,目标不可达。
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 无检查地配到另一条路径上。