规划目标
一次调用固定 FLOPs 与 read/write complexity 的权重,并让候选排名、改进阶段和最终选择使用同一目标。
本页目录
完整路径的评分公式
公开 PlannerObjective 比较完整路径的总 FLOPs F(P) 与 total read/write complexity R(P)。默认权重为 flops_weight=1、read_write_weight=64。
J(P) = wF F(P) + wR R(P)
- F(P)
- 路径的总 FLOPs
- R(P)
- 路径的总 read/write complexity
- wF、wR
- 本次调用指定的权重
使用同一组权重评价候选路径,选择 J(P) 最小的一条。
| 权重 | 目标标签 | 比较内容 |
|---|---|---|
| (1, 64) | flops_read_write | 默认加权目标 |
| (1, 0) | total_flops | 只比较总 FLOPs |
| (0, 1) | total_read_write | 只比较 total read/write complexity |
权重的含义
权重必须有限、非负,并且不能同时为 0。权重 64 是规划模型中的相对系数,不是字节数、缓存行大小或机器实测带宽。
一次调用只使用一套目标
权重在寻路入口确定。各生成器可以采用不同的局部启发式,但完整候选、树改进和最终选择使用同一个 PlannerObjective。
| 阶段 | objective 的作用 |
|---|---|
| random_greedy | 比较多个完整 trial 的结果 |
| bisect | 比较多个完整递归二分 trial 的结果 |
| reconfigure_path | 只接受 objective 严格改善的局部重构结果 |
| anneal_path(s) / temper_path(s) | 比较搜索中的完整树状态与最终结果 |
| 切片候选 | 在满足硬约束的方案之间比较总工作量 |
候选生成规则与顶层目标
random-greedy 的局部代价、图划分目标或退火中的 proposal 机制负责生成候选,并非新的顶层 objective。生成完整候选路径后,仍按本次调用的
PlannerObjective排名。
在 Python 中指定权重
from arctn import arctn_schedule
# 默认:FLOPs + 64 * read/write complexity
default_report = arctn_schedule(
inputs, output, size_dict,
flops_weight=1.0, read_write_weight=64.0,
)
# 纯 FLOPs
flops_report = arctn_schedule(
inputs, output, size_dict,
flops_weight=1.0, read_write_weight=0.0,
)
改变权重可能改变完整候选的排名、局部改进的接受结果和最终返回路径。搜索所用 objective 才是该次运行实际优化的目标。
| 报告字段 | 含义 |
|---|---|
| planner_objective | 本次调用使用的目标标签;权重分别由 flops_weight 和 read_write_weight 字段报告 |
| planner_objective_score_log2 | 最终返回计划加权目标的 log2;启用切片时计入全部切片 |
| log10_flops / log2_read_write | 返回路径在未切片网络上的指标;启用切片后,全部切片的总量分别读取 sliced_log10_flops_total 和 planner_log2_read_write |
目标、约束与报告指标
| 概念 | 回答的问题 | ArcTN 中的例子 |
|---|---|---|
| 规划目标 | 多个合法方案中哪一个更好 | F + 64R 或调用者指定的非负权重 |
| 硬约束 | 一个方案是否允许返回 | target_size 对单片二元结果的上限;单张量网络检查最终输出 |
| 生成方法 | 候选从哪里来 | random-greedy、bisection |
| 报告指标 | 路径还有哪些结构特征 | log2_peak_size、log2_total_size |
指定切片腿后,切片数量是这些腿维度的乘积。ArcTN 对切片方案的结构目标按“单片工作量 × slice 数量”比较,因此在 log2 域给单片 objective 加上 log2_n_slices。
n_slices = product(dim(leg) for leg in sliced_legs)
J_sliced = n_slices * J(one_slice)
log2(J_sliced) = log2(J(one_slice)) + log2(n_slices)
Warning
硬约束优先
target_size 是接口所检查结果大小的硬上限;只有满足限制的方案进入最终 objective 比较,软惩罚不能替代最终可行性复核。