本设计在 HyperParallel Trainer 中打通 DeepSeek-V3 16B 的 MXFP8/HiF8 在线低精预训练, 目标运行环境为 Ascend A5 + PyTorch/torch_npu。两种低精格式复用同一套模型加工、 forward/dgrad/wgrad、FSDP2、Optimizer 和 DCP 生命周期,只在量化策略、物理存储和 NPU 算子参数上分开实现。
本期支持:
DeepseekV3Config
TP=1、CP=1、EP=1、PP=1
mxfp8_e4m3 + mx_block
hif8 + current
本期不支持:
F.linear/mm
不提供 BF16/FP32 计算 fallback。fallback_to_unsharded 只允许不满足低精分片约束的参数 保持复制,参数对应的 MM/GMM 仍使用所选低精格式。
fallback_to_unsharded
低精组件范围同时参考论文边界和开源工程实现:
TELinear
TEColumnParallelLinear
TERowParallelLinear
linear_proj
fp8_dot_product_attention
fp8_multi_head_attention
本设计对 MXFP8 和 HiF8 使用相同的模块边界:MLA projection 作为普通 GEMM 进入低精; 核心 attention、RoPE 和 Q/K normalization 保持 BF16/FP32。目标必须按完整 FQN 显式列出, 不使用 self_attn.*_proj 或全模型 *_proj 通配。
self_attn.*_proj
*_proj
gate_proj/up_proj/down_proj
nn.Linear → LowPrecisionLinear
gate_up_proj/down_proj
q/kv
o_proj
nn.Linear
flowchart LR YAML[Trainer YAML] --> PLAN[Model Plan / Replacement Plan] PLAN --> DENSE[LowPrecisionLinear] PLAN --> DSV3[DeepseekV3LowPrecisionExperts] DENSE --> FUNC[Low-precision Functional] DSV3 --> FUNC FUNC --> RECIPE[Recipe / Role Quantizers] FUNC --> OPS[NPU MM/GMM Ops] RECIPE --> QT[MXFP8Tensor / HiF8Tensor] QT --> LIFE[FSDP / Optimizer / DCP Adapter]
依赖只能沿上图向下。DeepSeek adapter 可以依赖低精通用模块,低精通用模块不得反向导入 DeepSeek;HyperParallel core 只识别通用扩展协议,不得依赖 MXFP8 或 HiF8 类型。
torch_npu.hifloat8
row_data/row_scale + col_data/col_scale
data + scale
MXFP8 对每个 32 元素 block 独立缩放。更细的 scale 粒度能够适应同一 Tensor 内不同区域的 数值范围,但 row/col 量化方向对应不同的物理 data/scale,权重需要同时保存两个方向:
MXFP8 weight ├── row_data + row_scale(E8M0) └── col_data + col_scale(E8M0)
FSDP forward 可以根据计算类型只收集一个方向;backward 同时覆盖 dgrad/wgrad,需要完整的 row 和 col 表示。分片边界还必须与 MX scale tile 对齐。
HiF8 current scaling 的 Dense Tensor 使用当前 Tensor 的一个 FP32 scale,Grouped Tensor 使用按 expert/group 划分的 scale。scale 不依赖 row/col block,因此同一份量化 data 可以用于 不同 GEMM 方向:
HiF8 weight ├── data(hifloat8) └── scale(FP32 scalar or [E])
FSDP forward/backward 均收集同一份 data+scale。Dense 权重在各 rank 从 FP32 master shard 刷新前,需要对 amax 做 FSDP-group MAX reduce,保证所有 shard 使用同一 scale;Grouped 权重的 scale 与完整 expert 一起分片和重建。
从格式特征推测,MXFP8 的细粒度 block scale 对局部数值分布差异适应更强,但物理存储、 量化和通信组织更复杂;HiF8 的权重表示和通信更简单,但 Dense per-tensor scale 更容易受 异常大值影响。该描述只是设计层面的预期,最终精度、吞吐和显存结论必须通过 A5 相同模型、 相同 seed 的对比测试确定。
两种格式不建立两套模型或 Trainer 流程,共用:
plan_overrides
LowPrecisionLinear
实现分叉只收敛在格式层:
LowPrecisionRecipe ├── MXFP8Quantizer → MXFP8Tensor → npu_mxfp8 └── HiF8Quantizer → HiF8Tensor → npu_hif8
HiF8 初始数值策略冻结为:
15.0
224.0
scale = valid(amax) ? amax / format_max : 1.0
这些数值不作为普通模型配置开放;它们属于经过算子与精度验证的 HiF8 Recipe。后续若硬件 契约变化,通过新增 recipe 版本演进,避免同一个 recipe 名称在不同环境中产生不同语义。
三类配置各自只有一个职责:
LowPrecisionConfig
LowPrecisionPolicy
ModuleReplacementPlan
新增 from_hf_config() 入口,只加载 HF 配置并随机初始化模型,不加载 DeepSeek-V3 671B checkpoint:
from_hf_config()
HyperAutoModelForCausalLM.from_hf_config( config_name_or_path, *, config_overrides=None, torch_dtype=torch.bfloat16, )
冻结的 16B 规格示例:
model: _target_: hyper_models._transformers.HyperAutoModelForCausalLM.from_hf_config config_name_or_path: <hf-deepseek-v3-config-path> config_overrides: vocab_size: 102400 hidden_size: 2048 intermediate_size: 10944 moe_intermediate_size: 1408 num_hidden_layers: 28 first_k_dense_replace: 1 num_attention_heads: 16 num_key_value_heads: 16 n_routed_experts: 64 n_shared_experts: 2 q_lora_rank: null kv_lora_rank: 3072 qk_nope_head_dim: 64 qk_rope_head_dim: 64 v_head_dim: 128 num_experts_per_tok: 6 n_group: 1 topk_group: 1 norm_topk_prob: true routed_scaling_factor: 2.5 max_position_embeddings: 4096 use_cache: false torch_dtype: bfloat16 attn_implementation: sdpa force_hf: true
模型构建后必须校验层数、Dense/MoE 排布、expert 数、投影 shape 和总参数量,不能让来源 HF 配置隐式改变 16B 规格。
LowPrecisionConfig 只回答“怎么算”,不包含模块 FQN:
@dataclass(frozen=True) class LowPrecisionConfig: enabled: bool = False format: Literal["mxfp8_e4m3", "hif8"] = "mxfp8_e4m3" scaling: Literal["mx_block", "current"] = "mx_block"
合法组合:
format
scaling
mxfp8_e4m3
mx_block
hif8
current
格式与 scaling 组合不合法时,配置解析阶段直接报错。HiF8 current scaling 固定使用当前 amax,不开放安全余量等调节参数。解析后构建只读 LowPrecisionPolicy,再由模型加工上下文 传给 replacement factory。低精模块不直接读取 Trainer YAML。
YAML 示例:
# 二选一:MXFP8 low_precision: enabled: true format: mxfp8_e4m3 scaling: mx_block
# 二选一:HiF8 current scaling low_precision: enabled: true format: hif8 scaling: current
模块目标不随格式变化,因此 replacement factory 使用格式无关名称:
fsdp_config: dp_shard_size: 8 reshard_after_forward: true comm_fusion: false fallback_to_unsharded: false plan_overrides: - match: - "model.layers.*.self_attn.q_proj" - "model.layers.*.self_attn.kv_a_proj_with_mqa" - "model.layers.*.self_attn.kv_b_proj" - "model.layers.*.self_attn.o_proj" when: low_precision module_type: torch.nn.Linear exact_type: true replace_module: _target_: hyper_models.components.training.low_precision.modules.linear.replace_low_precision_linear - match: - "model.layers.0.mlp.*_proj" - "model.layers.*.mlp.shared_experts.*_proj" when: low_precision module_type: torch.nn.Linear exact_type: true replace_module: _target_: hyper_models.components.training.low_precision.modules.linear.replace_low_precision_linear - match: "model.layers.*.mlp.experts" when: low_precision module_type: transformers.models.deepseek_v3.modeling_deepseek_v3.DeepseekV3Experts exact_type: true replace_module: _target_: hyper_models.components.models.adapters.deepseek_v3.low_precision.replace_deepseek_v3_low_precision_experts
配置约束:
q_lora_rank=null
q_proj
q_a_proj/q_b_proj
exact_type=true
type(module) is configured_type
enabled=false
torch_npu
构建阶段的交接关系:
HF config + overrides → meta device 构建原始 DeepSeek-V3 模块树 │ plan_overrides ─┴→ 编译 ModuleReplacementPlan(只读取结构) │ └→ 按 HF 语义物化并初始化高精 weight/bias LowPrecisionConfig → LowPrecisionPolicy / Recipe ─────┐ ModuleReplacementPlan ─────────────────┤ ↓ → LowPrecisionLinear.from_linear(..., policy) → DeepseekV3LowPrecisionExperts(..., policy) → finalize 高精 weight 为对应格式 QuantizedWeightTensor → FSDP2 fully_shard
低精替换复用 Model Plan/Apply,不保留独立 Converter:
完整 model.named_modules() 扫描 → match、module_type、exact_type 校验 → 按模块 identity 合并 alias FQN → 生成不可变 Replacement Plan → 原模型完成高精参数初始化 → factory 构造并校验全部 replacement → 原子安装 replacement → 在稳定的新模块树上执行 Sharding Plan 和 FSDP2
replacement 必须保持模块 FQN、Parameter/Buffer 名称、共享关系、requires_grad、checkpoint key 和 forward 签名。它允许改变具体模块类型并增加 policy、recipe 和 quantizer 等普通 Python 属性。
requires_grad
只将 YAML 显式选中的 exact nn.Linear 替换为 LowPrecisionLinear:
class LowPrecisionLinear(nn.Linear): policy: LowPrecisionPolicy recipe: LowPrecisionRecipe weight: QuantizedWeightTensor def forward(self, inputs): ... def refresh_weight_storage(self, master_shard): ...
from_linear() 复用原始 weight/bias,不重新初始化或复制 Parameter。高精权重物化后, finalize 根据 policy 转换成 MXFP8Tensor 或 HiF8Tensor。bias 始终保持 BF16/FP32。
from_linear()
weight/bias
MXFP8Tensor
HiF8Tensor
HF DeepseekV3Experts 已持有堆叠的 3D 参数:
DeepseekV3Experts
gate_up_proj: [E, 2F, H] down_proj: [E, H, F]
因此整体替换为 DeepseekV3LowPrecisionExperts,不递归查找内部 Linear,也不再次堆叠 expert。该模块只适配 DeepSeek 参数名、forward 签名和 top-k 聚合;MXFP8/HiF8 Grouped SwiGLU 的计算与 autograd 放在通用 functional/grouped_linear.py。
DeepseekV3LowPrecisionExperts
functional/grouped_linear.py
shared experts 不是 routed expert container,其内部 exact Linear 继续走 Dense replacement。
hyper_models/ ├── trainer/config.py ├── _transformers/infrastructure.py └── components/ ├── distributed/ │ ├── module_replacement.py │ └── injection.py ├── models/adapters/deepseek_v3/ │ └── low_precision.py └── training/low_precision/ ├── config.py ├── policy.py ├── recipe.py ├── tensor/ │ ├── base.py │ ├── operand.py │ ├── mxfp8_tensor.py │ └── hif8_tensor.py ├── quantizers/ │ ├── base.py │ ├── mxfp8.py │ └── hif8.py ├── ops/ │ ├── npu_mxfp8.py │ └── npu_hif8.py ├── functional/ │ ├── linear.py │ └── grouped_linear.py ├── modules/ │ └── linear.py └── integration/ ├── fsdp_adapter.py ├── optimizer_adapter.py └── checkpoint_adapter.py
distributed/module_replacement.py
modules/linear.py
models/adapters/deepseek_v3/low_precision.py
recipe.py
functional/linear.py
quantizers/mxfp8.py
quantizers/hif8.py
ops/npu_mxfp8.py
ops/npu_hif8.py
integration/*_adapter.py
首期不增加通用 modules/grouped_linear.py:PyTorch/HF 没有统一的 Grouped Experts Module 接口。通用 GMM 放在 functional,DeepSeek 的有状态模块放在模型 adapter。
modules/grouped_linear.py
LowPrecisionRecipe 按角色创建 Quantizer,而不是由 Linear 在运行时判断格式:
LowPrecisionRecipe
class QuantizationRole(str, Enum): INPUT_FWD = "input_fwd" WEIGHT_FWD = "weight_fwd" GRAD_OUTPUT = "grad_output" INPUT_BWD = "input_bwd" WEIGHT_BWD = "weight_bwd"
MXFP8 可让多个角色共享同一个无状态 Quantizer。HiF8 current scaling 中 input/weight 的 format max 相同,可以共享不可变配置;gradient 使用不同 format max,使用独立 Quantizer。
“角色”和“方向”是两个概念:角色决定 dtype、format max 和 scaling 状态;row/col 决定当前 MM/GMM 的物理布局。
MXFP8Tensor 是跨 micro-batch 持续存在的权重 Owner:
逻辑 shape / dtype / requires_grad ├── row_data ├── row_scale E8M0 ├── col_data └── col_scale E8M0
四份物理存储由 FP32 master 派生,不进入 checkpoint。
HiF8 per-tensor/per-group scale 与 row/col block 无关,因此不复制两份相同 data:
逻辑 shape / dtype / requires_grad ├── data hifloat8 └── scale FP32 scalar 或 [E]
Dense 权重使用标量 scale;3D routed-expert 权重使用 [E] scale。select(layout) 通过视图和 transpose 元数据构造 operand,不重新量化或复制 data。
[E]
select(layout)
QuantizedOperand
QuantizedOperand 表示一次 MM/GMM 使用的短生命周期输入:
@dataclass(frozen=True) class QuantizedOperand: data: torch.Tensor scale: torch.Tensor logical_dtype: torch.dtype format: LowPrecisionFormat orientation: Literal["row", "col"]
activation 量化结果和权重 select() 均返回 operand。FSDP 只处理持久化 QuantizedWeightTensor,不处理 activation operand。
select()
QuantizedWeightTensor
module.forward → functional autograd ├── 根据 QuantizationRole 取得 Quantizer ├── activation/gradient → quantize → QuantizedOperand ├── QuantizedWeightTensor.select(row/col) → QuantizedOperand └── format ops → BF16 output
INPUT_FWD
WEIGHT_FWD
GRAD_OUTPUT
WEIGHT_BWD
INPUT_BWD
方向表是 GEMM/GMM 契约,对 MXFP8 决定选择哪份 physical data;对 HiF8 决定 data 的逻辑视图 和算子 transpose 参数。
activation/gradient 每次动态执行 MX block 量化。权重在 optimizer step 后从 FP32 master 一次性生成 row/col data/scale,在后续多个 micro-batch 和 activation recompute 中复用。
每次 activation/gradient 量化都从当前 Tensor 计算 amax 和 FP32 scale。持久化权重只在 optimizer step 后重新计算 scale 和 data,不在每个 forward 重复量化。
HyperParallel core 只提供格式无关协议:
@dataclass(frozen=True) class FSDPGatherContext: phase: Literal["forward", "backward"] reshard_after_forward: bool param_fqn: str class FSDPLocalTensorExtension(Protocol): def fsdp_pre_all_gather(self, context): """返回本阶段需要通信的 physical tensors 和 metadata。""" def fsdp_post_all_gather(self, outputs, metadata, out=None): """由通信输出重建完整逻辑 Tensor。"""
FSDP 不理解 row/col、format 或 scale,只分别 all-gather pre 返回的物理 Tensor,并把 metadata 原样传回 post。
pre
post
reshard_after_forward
true
false
reshard_after_forward=false 时,首次 all-gather 直接准备两个方向。backward 同时覆盖 dgrad 和 wgrad 所需数据,也准备两个方向。
reshard_after_forward=false
MXFP8 分片合法性:
[M, N]
M_local
M_local % 64 == 0
[E_local, D1, D2]
(E_local * D1) % 64 == 0
fallback_to_unsharded=false
HiF8 权重只有一份 data/scale,不按 forward/backward 选择不同方向:
data_shard + replicated scalar scale
data[E_local,...] + scale[E_local]
scale[E]
2D Dense 的每个 FSDP rank 必须使用相同 scale。刷新权重前先对本地 master-shard amax 在 FSDP group 做 MAX reduce,再以同一 scale 分别量化各 data shard;这样 all-gather 后与使用 同一 scale 量化完整权重数学等价。post 校验各 rank gather 到的 scalar scale 一致后折叠为 一个标量。
3D Grouped 参数只允许沿 expert 维切分,并保持单个 expert 完整;scale 随 expert shard 一起 all-gather。若 FSDP 布局会切开单个 expert,则直接报错或由 fallback_to_unsharded 保持复制。
本地低精 shard 在整个梯度累积周期内常驻;all-gather 得到的完整物理 Tensor 只在当前 unshard 窗口存在,由 FSDP 按 reshard 策略释放。二者不能互相释放。
训练状态:
FP32 master shard optimizer 更新的权威权重 FP32 exp_avg/exp_avg_sq optimizer 状态 低精 weight shard 由 master 派生的计算存储
首次构建时不能从低精权重反量化得到 FP32 master。高精参数 finalize 时临时保留初始化值; FSDP 完成本地分片后,用对应高精 shard 创建 FP32 master,随后释放初始化临时值。
统一 step 数据流:
梯度累积完成 → FSDP reduce-scatter 得到本地 BF16/FP32 gradient → optimizer 使用 FP32 master 和 FP32 状态更新 → 从新 master 刷新格式专属低精权重 MXFP8: dual-axis data + E8M0 scales HiF8 current: current amax + FP32 scale + data → 下一 optimizer step 内复用
刷新只发生在 optimizer.step() 成功后,不公开“已失效但尚未重建”状态。zero_grad() 和 activation recompute 不刷新权重。若 optimizer step 被 overflow/异常跳过,则不刷新权重。
optimizer.step()
zero_grad()
仅支持 DCP:
exp_avg/exp_avg_sq
1. 读取 HF config,覆盖并校验冻结的 16B 规格 2. 在 meta device 构建原始 HF 模型 3. 编译 ModuleReplacementPlan 4. 按原 HF 语义物化并初始化高精参数 5. 应用 Dense/Experts replacement 6. 根据 policy finalize MXFP8Tensor 或 HiF8Tensor 7. 执行 Sharding Plan 和 FSDP2 fully_shard 8. 创建 FP32 master 与 FP32 optimizer state
forward/recompute → 按角色量化 activation → unshard 低精 weight → MXFP8/HiF8 MM/GMM backward → 按角色量化 grad_output/input → MXFP8/HiF8 dgrad/wgrad → FSDP reduce-scatter gradient optimizer.step → 更新 FP32 master/state → 刷新低精 weight shard checkpoint(如到保存步) → 保存逻辑模型和 optimizer
任何读取真实数据、调用 NPU kernel 或创建通信组的动作,都不得发生在 meta 构建和 replacement 编译阶段。
本节冻结 Q3 转测依赖的接口。以下是语义契约,最终 Python 包路径可随目录实现调整,但调用方 不应依赖具体格式类的私有字段。
输入来自 Trainer YAML;输出为不可变 LowPrecisionPolicy。配置层只允许两个组合:
mxfp8_e4m3 + mx_block hif8 + current
其他组合在模型构建前报错。enabled=false 时不得导入 torch_npu、替换模块或创建低精状态。
def replace_low_precision_linear( source: nn.Linear, *, fqn: str, policy: LowPrecisionPolicy, ) -> LowPrecisionLinear: ... def replace_deepseek_v3_low_precision_experts( source: nn.Module, *, fqn: str, policy: LowPrecisionPolicy, ) -> DeepseekV3LowPrecisionExperts: ...
调用方是 ModuleReplacementPlan。factory 必须复用原 Parameter/Buffer,不得改变 FQN、共享 关系、requires_grad、state_dict key 和 forward 输入输出契约;失败时 Plan 不得部分修改模型。
class LowPrecisionRecipe(Protocol): def quantizer(self, role: QuantizationRole) -> Quantizer: ... class Quantizer(Protocol): def quantize( self, tensor: torch.Tensor, *, orientation: Literal["row", "col"], group_list: torch.Tensor | None = None, ) -> QuantizedOperand: ...
QuantizationRole 决定 input/weight/gradient 的格式参数,orientation 决定 GEMM/GMM layout。 Quantizer 不遍历模型、不执行 collective,也不负责 optimizer 或 checkpoint。
QuantizationRole
orientation
class QuantizedWeightTensor(torch.Tensor): def select(self, orientation: Literal["row", "col"]) -> QuantizedOperand: ... def refresh_from_master(self, master_shard: torch.Tensor) -> None: ... def low_precision_matmul( left: QuantizedOperand, right: QuantizedOperand, *, layout: Literal["NN", "NT", "TN"], output_dtype: torch.dtype, ) -> torch.Tensor: ... def low_precision_grouped_matmul( left: QuantizedOperand, right: QuantizedOperand, group_list: torch.Tensor, *, layout: Literal["NN", "NT", "TN"], output_dtype: torch.dtype, ) -> torch.Tensor: ...
公开计算接口使用 QuantizedOperand,不让上层直接拼装 data/scale。格式专属 NPU adapter 负责 dtype、scale dtype、transpose 和算子参数校验,任何失败都不得退回高精 GEMM。
FSDP2 只调用第 9 章的 fsdp_pre_all_gather(context) 和 fsdp_post_all_gather(outputs, metadata, out)。扩展返回物理 Tensor 的有序 tuple 和可序列化 metadata;post 必须校验输出数量、shape、dtype 和顺序。普通 Tensor 不进入该接口。
fsdp_pre_all_gather(context)
fsdp_post_all_gather(outputs, metadata, out)
Optimizer adapter 只在成功更新 FP32 master 后调用:
def refresh_low_precision_weights( params: Iterable[QuantizedWeightTensor], master_params: Iterable[torch.Tensor], ) -> None: ...
DCP adapter 对外提供高精 checkpoint 视图:模型为 BF16 逻辑权重,optimizer 为 FP32 master、 FP32 moments 和 step。MXFP8/HiF8 data/scale 不属于持久化接口。
from_hf_config
必须保持的加工顺序:
HF/meta 构建 → 高精参数物化 → Module Replacement Plan → 低精权重 finalize → Sharding Plan / FSDP2 → activation checkpoint wrapper → optimizer 与 DCP 注册
顺序改变可能导致 meta 阶段调用 NPU、FSDP 看不到低精 Tensor、hook 重复注册或 optimizer 拿不到高精 master,属于启动时必须检查的集成错误。
启动时按 policy 校验,不通过直接失败:
TP/CP/EP/PP > 1
运行阶段不得因 capability、shape 或算子不支持退回 BF16/FP32 GEMM。
完整限制矩阵:
UT 默认在 CPU/mock NPU Ops 上执行;验证 Python 契约、数据流和失败语义。真实 kernel 数值只在 A5 ST/转测验证,UT 不伪装成硬件验收。
format/scaling
amax/format_max
select(row/col)
Quantizer 数值 UT 使用确定性小 Tensor 手算 scale 和反量化误差;MM/GMM Python UT 使用 fake operator 记录调用参数,不把 fake 输出作为真实低精精度结论。
NN/NT/TN
reshard_after_forward=true/false
comm_fusion=true
activation_checkpoint=off/full
ST 必须经过生产 Trainer 入口。STProbe 只观测模块类型、kernel、FSDP physical tensors、scale 和 refresh 次数,不在测试中手工执行 replacement、量化、all-gather 或权重刷新。
STProbe
# 伪代码:接口名称表达测试职责,不要求与最终实现逐字一致。 @dataclass(frozen=True) class STCase: name: str low_precision: dict | None # None 表示 BF16 baseline reshard_after_forward: bool activation_checkpoint: str # "off" / "full" grad_accum_steps: int CASES = [ STCase("bf16", None, True, "off", 1), STCase( "mxfp8", {"enabled": True, "format": "mxfp8_e4m3", "scaling": "mx_block"}, True, "off", 1, ), STCase( "hif8_current", {"enabled": True, "format": "hif8", "scaling": "current"}, True, "off", 1, ), ] def run_st(case: STCase, checkpoint_dir: Path) -> STResult: init_process_group(backend="hccl") set_deterministic_seed(2026) # 所有 case 使用相同 seed 和 batch 顺序 cfg = load_yaml("dsv3_16b_fsdp2_st.yaml") cfg.low_precision = case.low_precision cfg.fsdp_config.reshard_after_forward = case.reshard_after_forward cfg.gradient_checkpointing.activation_checkpoint = case.activation_checkpoint cfg.gradient_accumulation_steps = case.grad_accum_steps # setup() 必须走正式顺序:HF 构建 → replacement → finalize → FSDP2 → optimizer。 trainer = Trainer() probe = STProbe(read_only=True) trainer.register_test_probe(probe) trainer.setup(cfg) assert_model_scope(trainer.model, case, probe) assert_optimizer_fp32_state(trainer.optimizer) if case.low_precision is not None: assert_no_high_precision_gemm_fallback(probe) step_records = [] for step in range(WARMUP_STEPS + MEASURE_STEPS): trainer.optimizer.zero_grad(set_to_none=True) probe.begin_optimizer_step(step) for micro_step in range(case.grad_accum_steps): batch = deterministic_batch(step, micro_step) loss = trainer.forward_backward( batch, loss_scale=1.0 / case.grad_accum_steps, ) assert torch.isfinite(loss) assert_all_gradients_finite(trainer.model) # 由生产 optimizer adapter 在成功 step 后刷新;ST 只记录事件。 updated = trainer.optimizer_step() probe.end_optimizer_step(updated=updated) if updated and case.low_precision is not None: probe.assert_one_weight_refresh_per_parameter(step) if not updated: probe.assert_no_weight_refresh(step) torch.npu.synchronize() step_records.append( collect_step_record( step=step, loss=loss, grad_norm=trainer.grad_norm, elapsed=probe.step_elapsed, peak_memory=torch.npu.max_memory_allocated(), ) ) if step == SAVE_STEP: trainer.save_dcp(checkpoint_dir) if case.low_precision is not None: assert_low_precision_kernel_contract(case, probe) assert_fsdp_physical_tensor_contract(case, probe) assert_scale_contract(case, probe) return STResult( case=case, records=step_records, checkpoint_dir=checkpoint_dir, trace=probe.export_trace(), )
关键断言伪代码:
def assert_low_precision_kernel_contract(case, probe): # Dense 与 routed experts 的三种反向角色必须全部真实执行。 expected = { ("dense", "fprop"), ("dense", "dgrad"), ("dense", "wgrad"), ("grouped", "fprop"), ("grouped", "dgrad"), ("grouped", "wgrad"), } assert expected <= probe.executed_kernel_roles() assert probe.high_precision_fallback_count == 0 if case.name == "mxfp8": probe.assert_all_kernel_format("mxfp8_e4m3") probe.assert_mxfp8_scale_dtype("e8m0") elif case.name == "hif8_current": probe.assert_all_kernel_format("hif8") probe.assert_hif8_scale_dtype(torch.float32) probe.assert_hif8_format_max(input_weight=15.0, gradient=224.0) def assert_fsdp_physical_tensor_contract(case, probe): for event in probe.fsdp_gather_events: event.assert_post_matches_pre_schema() if case.name == "mxfp8": probe.assert_mxfp8_forward_direction_by_compute_kind() probe.assert_mxfp8_backward_has_row_and_col() elif case.name == "hif8_current": probe.assert_hif8_forward_backward_use_data_and_scale() probe.assert_hif8_dense_scale_equal_across_ranks() probe.assert_hif8_grouped_scale_matches_expert_order() def assert_scale_contract(case, probe): if case.name != "hif8_current": return for event in probe.hif8_quant_events: expected_max = 224.0 if event.role == "grad_output" else 15.0 expected_scale = safe_scale(event.current_amax, expected_max) torch.testing.assert_close(event.scale, expected_scale) assert event.scale.dtype == torch.float32 assert torch.isfinite(event.scale).all() and (event.scale > 0).all()
DCP 续训使用新的 Trainer 实例,防止旧进程内对象掩盖恢复问题:
def run_resume_st(case, continuous_result): resumed = Trainer() resumed.setup(load_same_case_config(case)) resumed.load_dcp(continuous_result.checkpoint_dir) assert resumed.global_step == SAVE_STEP assert_fp32_master_and_moments(resumed.optimizer) assert_checkpoint_has_no_derived_low_precision_storage( continuous_result.checkpoint_dir ) # load_dcp 后由生产 adapter 从 FP32 master 重建低精权重。 assert_rebuilt_weight_storage(resumed.model, case) for step in range(SAVE_STEP, SAVE_STEP + RESUME_STEPS): batch = deterministic_batch(step, 0) resumed_record = resumed.train_step(batch) continuous_record = continuous_result.record(step) assert_resume_metrics_close(resumed_record, continuous_record)
性能汇总伪代码:
bf16 = run_st(BF16_CASE, tmp_path / "bf16") mxfp8 = run_st(MXFP8_CASE, tmp_path / "mxfp8") hif8 = run_st(HIF8_CURRENT_CASE, tmp_path / "hif8") for low_precision in (mxfp8, hif8): # 丢弃 warmup,只统计相同 token 数的连续 optimizer steps。 assert low_precision.steady_tokens_per_second >= bf16.steady_tokens_per_second assert low_precision.average_step_time <= bf16.average_step_time
8 卡启动形式:
torchrun --nproc_per_node=8 tests/st/low_precision/test_dsv3_16b_fsdp2.py \ --config configs/st/dsv3_16b_fsdp2.yaml \ --case hif8_current
Q3 转测使用单机 8 卡 Ascend A5(每卡 96 GiB)、PyTorch + torch_npu,固定 DeepSeek-V3 16B 规格和同一份确定性训练数据。每次测试记录代码 commit、CANN/torch_npu/PyTorch 版本、 完整 YAML、随机种子和算子 trace。若 16B 因环境资源问题无法启动,只能用缩小模型定位,不能 以缩小模型代替最终验收。
固定配置维度:
format: BF16 baseline / MXFP8 / HiF8 current reshard_after_forward: true / false activation_checkpoint: off / full fallback_to_unsharded: false(主路径)/ true(专项) gradient_accumulation_steps: 1 / >1
activation_checkpoint=full
这里的“约定容差”必须在首轮 A5 BF16、MXFP8、HiF8 固定 seed 数据采集后,由算法负责人给出 具体数值并写回转测用例;在阈值冻结前,Q3-09 不得仅凭“曲线看起来接近”判定通过。
Q3 HiF8 专项验收仅覆盖 hif8 + current,不覆盖 delayed scaling。转测不需要构造、检查或 恢复 amax history、pending amax、scale update interval 等 delayed 状态。
[E_local,...]
[E_local]
HiF8 delayed scaling 配置必须在启动阶段报“不支持”,而不是被忽略、映射为 current 或回退到 高精计算。delayed scaling 的实现与验收统一放到第 18 章后续演进。
端到端 loss 不能单独证明 Fprop/Dgrad/Wgrad 正确,必须分层验收:
Q3 首轮允许“采集阈值”与功能转测并行,但发布前必须冻结上述指标的数值门槛。文档不预设 未经 A5 实测的统一误差阈值,避免用任意阈值掩盖格式和算子差异。
性能是 Q3 阻塞验收项。比较时必须固定模型、训练数据、sequence length、micro/global batch、 梯度累积、FSDP、reshard 和 recompute 配置;不得通过缩小 batch、减少有效 token 或改变并行 策略得到表面加速。关闭 profiling 和一次性初始化影响,预热后至少采集 20 个连续 optimizer step,以相同统计方法计算:
throughput_ratio = low_precision_tokens_per_second / bf16_tokens_per_second step_time_ratio = low_precision_average_step_time / bf16_average_step_time 通过条件: throughput_ratio >= 1.0 step_time_ratio <= 1.0
MXFP8 和 HiF8 current 必须分别满足上述条件,不能用一种格式的收益抵消另一种格式的劣化。 同时必须产出:
任一低精格式吞吐低于 BF16 或平均 step time 高于 BF16,均判定 Q3 性能验收失败。峰值显存 本期要求记录并解释,暂不设置相对 BF16 的阻塞阈值;若发现错误 fallback、重复量化、重复 all-gather 或缓存未释放,则同时按功能问题阻塞。
HiF8 delayed scaling 后续单独设计和交付。届时需要新增角色级 amax history、跨 rank amax 归并、optimizer-step 更新时序和 DCP scaling-state 保存恢复;本期不提前引入这些状态、配置或 生命周期接口。
需要补充说明Q3转测的范围、与其他模块相关性、UT设计,最好有个使用样例
DeepSeek-V3 16B FSDP2 MXFP8/HiF8 训练设计
1. 目标与范围
本设计在 HyperParallel Trainer 中打通 DeepSeek-V3 16B 的 MXFP8/HiF8 在线低精预训练,
目标运行环境为 Ascend A5 + PyTorch/torch_npu。两种低精格式复用同一套模型加工、
forward/dgrad/wgrad、FSDP2、Optimizer 和 DCP 生命周期,只在量化策略、物理存储和 NPU
算子参数上分开实现。
本期支持:
DeepseekV3Config构建冻结的 16B 结构并随机初始化;TP=1、CP=1、EP=1、PP=1,使用 FSDP2 做数据并行分片;forward/dgrad/wgrad 低精计算;
mxfp8_e4m3 + mx_block;hif8 + currentscaling;本期不支持:
F.linear/mm;不提供 BF16/FP32 计算 fallback。
fallback_to_unsharded只允许不满足低精分片约束的参数保持复制,参数对应的 MM/GMM 仍使用所选低精格式。
1.1 模型低精边界
低精组件范围同时参考论文边界和开源工程实现:
将主要 GEMM 的 Fprop、Dgrad 和 Wgrad 放到 FP8,同时将 embedding、output head、MoE
gate、normalization 和 attention operators 保持高精。论文没有逐个枚举 MLA projection,
因而不能仅凭该表判断 projection Linear 是否属于这里的 attention operators。
MLA layer spec
将 Q、KV 的 down/up projection 和 output projection 都交给 backend Linear 构建;其
Transformer Engine backend
对应
TELinear、TEColumnParallelLinear和TERowParallelLinear。MLA 实现还明确处理linear_proj保存量化输入,说明 projection Linear 位于 FP8 路径。
fp8_dot_product_attention和fp8_multi_head_attention均默认关闭。因此其默认边界是“projection Linear 可低精,QK、Softmax、PV 等核心
attention 计算保持高精”。这是工程实现证据,不反向证明 DeepSeek-V3 原始训练对每个
MLA projection 都使用了 FP8。
本设计对 MXFP8 和 HiF8 使用相同的模块边界:MLA projection 作为普通 GEMM 进入低精;
核心 attention、RoPE 和 Q/K normalization 保持 BF16/FP32。目标必须按完整 FQN 显式列出,
不使用
self_attn.*_proj或全模型*_proj通配。gate_proj/up_proj/down_projnn.Linear → LowPrecisionLineargate_proj/up_proj/down_projnn.Linear → LowPrecisionLineargate_up_proj/down_projq/kvdown/up projection、o_projnn.Linear2. 总体架构
flowchart LR YAML[Trainer YAML] --> PLAN[Model Plan / Replacement Plan] PLAN --> DENSE[LowPrecisionLinear] PLAN --> DSV3[DeepseekV3LowPrecisionExperts] DENSE --> FUNC[Low-precision Functional] DSV3 --> FUNC FUNC --> RECIPE[Recipe / Role Quantizers] FUNC --> OPS[NPU MM/GMM Ops] RECIPE --> QT[MXFP8Tensor / HiF8Tensor] QT --> LIFE[FSDP / Optimizer / DCP Adapter]依赖只能沿上图向下。DeepSeek adapter 可以依赖低精通用模块,低精通用模块不得反向导入
DeepSeek;HyperParallel core 只识别通用扩展协议,不得依赖 MXFP8 或 HiF8 类型。
3. 两种格式的统一与差异
torch_npu.hifloat8row_data/row_scale + col_data/col_scaledata + scale3.1 核心差别
MXFP8 对每个 32 元素 block 独立缩放。更细的 scale 粒度能够适应同一 Tensor 内不同区域的
数值范围,但 row/col 量化方向对应不同的物理 data/scale,权重需要同时保存两个方向:
FSDP forward 可以根据计算类型只收集一个方向;backward 同时覆盖 dgrad/wgrad,需要完整的
row 和 col 表示。分片边界还必须与 MX scale tile 对齐。
HiF8 current scaling 的 Dense Tensor 使用当前 Tensor 的一个 FP32 scale,Grouped Tensor
使用按 expert/group 划分的 scale。scale 不依赖 row/col block,因此同一份量化 data 可以用于
不同 GEMM 方向:
FSDP forward/backward 均收集同一份 data+scale。Dense 权重在各 rank 从 FP32 master shard
刷新前,需要对 amax 做 FSDP-group MAX reduce,保证所有 shard 使用同一 scale;Grouped
权重的 scale 与完整 expert 一起分片和重建。
从格式特征推测,MXFP8 的细粒度 block scale 对局部数值分布差异适应更强,但物理存储、
量化和通信组织更复杂;HiF8 的权重表示和通信更简单,但 Dense per-tensor scale 更容易受
异常大值影响。该描述只是设计层面的预期,最终精度、吞吐和显存结论必须通过 A5 相同模型、
相同 seed 的对比测试确定。
3.2 共用训练流程
两种格式不建立两套模型或 Trainer 流程,共用:
plan_overrides和 Module Replacement Plan;LowPrecisionLinear与 DeepSeek Experts adapter;实现分叉只收敛在格式层:
HiF8 初始数值策略冻结为:
15.0;224.0;scale = valid(amax) ? amax / format_max : 1.0;这些数值不作为普通模型配置开放;它们属于经过算子与精度验证的 HiF8 Recipe。后续若硬件
契约变化,通过新增 recipe 版本演进,避免同一个 recipe 名称在不同环境中产生不同语义。
4. 配置入口
三类配置各自只有一个职责:
LowPrecisionConfigLowPrecisionPolicyplan_overridesModuleReplacementPlan4.1 HF 模型配置
新增
from_hf_config()入口,只加载 HF 配置并随机初始化模型,不加载 DeepSeek-V3 671Bcheckpoint:
HyperAutoModelForCausalLM.from_hf_config( config_name_or_path, *, config_overrides=None, torch_dtype=torch.bfloat16, )冻结的 16B 规格示例:
model: _target_: hyper_models._transformers.HyperAutoModelForCausalLM.from_hf_config config_name_or_path: <hf-deepseek-v3-config-path> config_overrides: vocab_size: 102400 hidden_size: 2048 intermediate_size: 10944 moe_intermediate_size: 1408 num_hidden_layers: 28 first_k_dense_replace: 1 num_attention_heads: 16 num_key_value_heads: 16 n_routed_experts: 64 n_shared_experts: 2 q_lora_rank: null kv_lora_rank: 3072 qk_nope_head_dim: 64 qk_rope_head_dim: 64 v_head_dim: 128 num_experts_per_tok: 6 n_group: 1 topk_group: 1 norm_topk_prob: true routed_scaling_factor: 2.5 max_position_embeddings: 4096 use_cache: false torch_dtype: bfloat16 attn_implementation: sdpa force_hf: true模型构建后必须校验层数、Dense/MoE 排布、expert 数、投影 shape 和总参数量,不能让来源
HF 配置隐式改变 16B 规格。
4.2 低精策略配置
LowPrecisionConfig只回答“怎么算”,不包含模块 FQN:@dataclass(frozen=True) class LowPrecisionConfig: enabled: bool = False format: Literal["mxfp8_e4m3", "hif8"] = "mxfp8_e4m3" scaling: Literal["mx_block", "current"] = "mx_block"合法组合:
formatscalingmxfp8_e4m3mx_blockhif8current格式与 scaling 组合不合法时,配置解析阶段直接报错。HiF8 current scaling 固定使用当前
amax,不开放安全余量等调节参数。解析后构建只读
LowPrecisionPolicy,再由模型加工上下文传给 replacement factory。低精模块不直接读取 Trainer YAML。
YAML 示例:
# 二选一:MXFP8 low_precision: enabled: true format: mxfp8_e4m3 scaling: mx_block# 二选一:HiF8 current scaling low_precision: enabled: true format: hif8 scaling: current4.3 目标模块与 FSDP 配置
模块目标不随格式变化,因此 replacement factory 使用格式无关名称:
fsdp_config: dp_shard_size: 8 reshard_after_forward: true comm_fusion: false fallback_to_unsharded: false plan_overrides: - match: - "model.layers.*.self_attn.q_proj" - "model.layers.*.self_attn.kv_a_proj_with_mqa" - "model.layers.*.self_attn.kv_b_proj" - "model.layers.*.self_attn.o_proj" when: low_precision module_type: torch.nn.Linear exact_type: true replace_module: _target_: hyper_models.components.training.low_precision.modules.linear.replace_low_precision_linear - match: - "model.layers.0.mlp.*_proj" - "model.layers.*.mlp.shared_experts.*_proj" when: low_precision module_type: torch.nn.Linear exact_type: true replace_module: _target_: hyper_models.components.training.low_precision.modules.linear.replace_low_precision_linear - match: "model.layers.*.mlp.experts" when: low_precision module_type: transformers.models.deepseek_v3.modeling_deepseek_v3.DeepseekV3Experts exact_type: true replace_module: _target_: hyper_models.components.models.adapters.deepseek_v3.low_precision.replace_deepseek_v3_low_precision_experts配置约束:
plan_overrides不重复配置 format/scaling;q_lora_rank=null只配置q_proj,非空时改配q_a_proj/q_b_proj;exact_type=true使用type(module) is configured_type;enabled=false时不调用 factory 和 NPU capability 检查;torch_npu。构建阶段的交接关系:
5. 模块替换
低精替换复用 Model Plan/Apply,不保留独立 Converter:
replacement 必须保持模块 FQN、Parameter/Buffer 名称、共享关系、
requires_grad、checkpoint key和 forward 签名。它允许改变具体模块类型并增加 policy、recipe 和 quantizer 等普通 Python
属性。
5.1 Dense 与 shared experts
只将 YAML 显式选中的 exact
nn.Linear替换为LowPrecisionLinear:class LowPrecisionLinear(nn.Linear): policy: LowPrecisionPolicy recipe: LowPrecisionRecipe weight: QuantizedWeightTensor def forward(self, inputs): ... def refresh_weight_storage(self, master_shard): ...from_linear()复用原始weight/bias,不重新初始化或复制 Parameter。高精权重物化后,finalize 根据 policy 转换成
MXFP8Tensor或HiF8Tensor。bias 始终保持 BF16/FP32。5.2 Routed experts
HF
DeepseekV3Experts已持有堆叠的 3D 参数:因此整体替换为
DeepseekV3LowPrecisionExperts,不递归查找内部 Linear,也不再次堆叠expert。该模块只适配 DeepSeek 参数名、forward 签名和 top-k 聚合;MXFP8/HiF8 Grouped
SwiGLU 的计算与 autograd 放在通用
functional/grouped_linear.py。shared experts 不是 routed expert container,其内部 exact Linear 继续走 Dense replacement。
6. 内部目录
distributed/module_replacement.pymodules/linear.pyLowPrecisionLinear与 Dense factorymodels/adapters/deepseek_v3/low_precision.pyDeepseekV3LowPrecisionExperts与模型 factoryrecipe.pyfunctional/linear.pyfunctional/grouped_linear.pyquantizers/mxfp8.pyquantizers/hif8.pyops/npu_mxfp8.pyops/npu_hif8.pyintegration/*_adapter.py首期不增加通用
modules/grouped_linear.py:PyTorch/HF 没有统一的 Grouped Experts Module接口。通用 GMM 放在 functional,DeepSeek 的有状态模块放在模型 adapter。
7. Recipe、角色与数据结构
7.1 计算角色
LowPrecisionRecipe按角色创建 Quantizer,而不是由 Linear 在运行时判断格式:class QuantizationRole(str, Enum): INPUT_FWD = "input_fwd" WEIGHT_FWD = "weight_fwd" GRAD_OUTPUT = "grad_output" INPUT_BWD = "input_bwd" WEIGHT_BWD = "weight_bwd"MXFP8 可让多个角色共享同一个无状态 Quantizer。HiF8 current scaling 中 input/weight 的
format max 相同,可以共享不可变配置;gradient 使用不同 format max,使用独立 Quantizer。
“角色”和“方向”是两个概念:角色决定 dtype、format max 和 scaling 状态;row/col 决定当前
MM/GMM 的物理布局。
7.2
MXFP8TensorMXFP8Tensor是跨 micro-batch 持续存在的权重 Owner:四份物理存储由 FP32 master 派生,不进入 checkpoint。
7.3
HiF8TensorHiF8 per-tensor/per-group scale 与 row/col block 无关,因此不复制两份相同 data:
Dense 权重使用标量 scale;3D routed-expert 权重使用
[E]scale。select(layout)通过视图和transpose 元数据构造 operand,不重新量化或复制 data。
7.4
QuantizedOperandQuantizedOperand表示一次 MM/GMM 使用的短生命周期输入:@dataclass(frozen=True) class QuantizedOperand: data: torch.Tensor scale: torch.Tensor logical_dtype: torch.dtype format: LowPrecisionFormat orientation: Literal["row", "col"]activation 量化结果和权重
select()均返回 operand。FSDP 只处理持久化QuantizedWeightTensor,不处理 activation operand。8. 计算数据流
INPUT_FWD、WEIGHT_FWDGRAD_OUTPUT、WEIGHT_BWDGRAD_OUTPUT、INPUT_BWD方向表是 GEMM/GMM 契约,对 MXFP8 决定选择哪份 physical data;对 HiF8 决定 data 的逻辑视图
和算子 transpose 参数。
8.1 MXFP8
activation/gradient 每次动态执行 MX block 量化。权重在 optimizer step 后从 FP32 master
一次性生成 row/col data/scale,在后续多个 micro-batch 和 activation recompute 中复用。
8.2 HiF8 current scaling
每次 activation/gradient 量化都从当前 Tensor 计算 amax 和 FP32 scale。持久化权重只在
optimizer step 后重新计算 scale 和 data,不在每个 forward 重复量化。
9. FSDP2 多存储扩展
HyperParallel core 只提供格式无关协议:
@dataclass(frozen=True) class FSDPGatherContext: phase: Literal["forward", "backward"] reshard_after_forward: bool param_fqn: str class FSDPLocalTensorExtension(Protocol): def fsdp_pre_all_gather(self, context): """返回本阶段需要通信的 physical tensors 和 metadata。""" def fsdp_post_all_gather(self, outputs, metadata, out=None): """由通信输出重建完整逻辑 Tensor。"""FSDP 不理解 row/col、format 或 scale,只分别 all-gather
pre返回的物理 Tensor,并把metadata 原样传回
post。9.1 MXFP8 gather
reshard_after_forwardtruefalsereshard_after_forward=false时,首次 all-gather 直接准备两个方向。backward 同时覆盖 dgrad和 wgrad 所需数据,也准备两个方向。
MXFP8 分片合法性:
[M, N]:本地M_local满足 scale tile 边界,当前要求M_local % 64 == 0;[E_local, D1, D2]:前两维为展平行维,要求(E_local * D1) % 64 == 0;fallback_to_unsharded=false时直接报错,为true时保持复制。9.2 HiF8 gather
HiF8 权重只有一份 data/scale,不按 forward/backward 选择不同方向:
data_shard + replicated scalar scaledata[E_local,...] + scale[E_local]scale[E]2D Dense 的每个 FSDP rank 必须使用相同 scale。刷新权重前先对本地 master-shard amax 在
FSDP group 做 MAX reduce,再以同一 scale 分别量化各 data shard;这样 all-gather 后与使用
同一 scale 量化完整权重数学等价。
post校验各 rank gather 到的 scalar scale 一致后折叠为一个标量。
3D Grouped 参数只允许沿 expert 维切分,并保持单个 expert 完整;scale 随 expert shard 一起
all-gather。若 FSDP 布局会切开单个 expert,则直接报错或由
fallback_to_unsharded保持复制。9.3 常驻与临时状态
本地低精 shard 在整个梯度累积周期内常驻;all-gather 得到的完整物理 Tensor 只在当前
unshard 窗口存在,由 FSDP 按 reshard 策略释放。二者不能互相释放。
10. Optimizer 生命周期
训练状态:
首次构建时不能从低精权重反量化得到 FP32 master。高精参数 finalize 时临时保留初始化值;
FSDP 完成本地分片后,用对应高精 shard 创建 FP32 master,随后释放初始化临时值。
统一 step 数据流:
刷新只发生在
optimizer.step()成功后,不公开“已失效但尚未重建”状态。zero_grad()和activation recompute 不刷新权重。若 optimizer step 被 overflow/异常跳过,则不刷新权重。
11. Checkpoint
仅支持 DCP:
exp_avg/exp_avg_sq;fallback_to_unsharded参数按普通复制参数处理。12. 构建与运行顺序
12.1 构建
12.2 单个训练 step
任何读取真实数据、调用 NPU kernel 或创建通信组的动作,都不得发生在 meta 构建和
replacement 编译阶段。
13. 对外接口设计
本节冻结 Q3 转测依赖的接口。以下是语义契约,最终 Python 包路径可随目录实现调整,但调用方
不应依赖具体格式类的私有字段。
13.1 Trainer 配置接口
@dataclass(frozen=True) class LowPrecisionConfig: enabled: bool = False format: Literal["mxfp8_e4m3", "hif8"] = "mxfp8_e4m3" scaling: Literal["mx_block", "current"] = "mx_block"输入来自 Trainer YAML;输出为不可变
LowPrecisionPolicy。配置层只允许两个组合:其他组合在模型构建前报错。
enabled=false时不得导入 torch_npu、替换模块或创建低精状态。13.2 模块替换接口
def replace_low_precision_linear( source: nn.Linear, *, fqn: str, policy: LowPrecisionPolicy, ) -> LowPrecisionLinear: ... def replace_deepseek_v3_low_precision_experts( source: nn.Module, *, fqn: str, policy: LowPrecisionPolicy, ) -> DeepseekV3LowPrecisionExperts: ...调用方是
ModuleReplacementPlan。factory 必须复用原 Parameter/Buffer,不得改变 FQN、共享关系、
requires_grad、state_dict key 和 forward 输入输出契约;失败时 Plan 不得部分修改模型。13.3 Recipe 与 Quantizer 接口
class LowPrecisionRecipe(Protocol): def quantizer(self, role: QuantizationRole) -> Quantizer: ... class Quantizer(Protocol): def quantize( self, tensor: torch.Tensor, *, orientation: Literal["row", "col"], group_list: torch.Tensor | None = None, ) -> QuantizedOperand: ...QuantizationRole决定 input/weight/gradient 的格式参数,orientation决定 GEMM/GMM layout。Quantizer 不遍历模型、不执行 collective,也不负责 optimizer 或 checkpoint。
13.4 低精 Tensor 与算子接口
class QuantizedWeightTensor(torch.Tensor): def select(self, orientation: Literal["row", "col"]) -> QuantizedOperand: ... def refresh_from_master(self, master_shard: torch.Tensor) -> None: ... def low_precision_matmul( left: QuantizedOperand, right: QuantizedOperand, *, layout: Literal["NN", "NT", "TN"], output_dtype: torch.dtype, ) -> torch.Tensor: ... def low_precision_grouped_matmul( left: QuantizedOperand, right: QuantizedOperand, group_list: torch.Tensor, *, layout: Literal["NN", "NT", "TN"], output_dtype: torch.dtype, ) -> torch.Tensor: ...公开计算接口使用
QuantizedOperand,不让上层直接拼装 data/scale。格式专属 NPU adapter负责 dtype、scale dtype、transpose 和算子参数校验,任何失败都不得退回高精 GEMM。
13.5 FSDP2 扩展接口
FSDP2 只调用第 9 章的
fsdp_pre_all_gather(context)和fsdp_post_all_gather(outputs, metadata, out)。扩展返回物理 Tensor 的有序 tuple 和可序列化metadata;
post必须校验输出数量、shape、dtype 和顺序。普通 Tensor 不进入该接口。13.6 Optimizer 与 DCP 接口
Optimizer adapter 只在成功更新 FP32 master 后调用:
def refresh_low_precision_weights( params: Iterable[QuantizedWeightTensor], master_params: Iterable[torch.Tensor], ) -> None: ...DCP adapter 对外提供高精 checkpoint 视图:模型为 BF16 逻辑权重,optimizer 为 FP32 master、
FP32 moments 和 step。MXFP8/HiF8 data/scale 不属于持久化接口。
14. 与其他模块的相关性
from_hf_config、16B override 和结构校验DeepseekV3LowPrecisionExperts必须保持的加工顺序:
顺序改变可能导致 meta 阶段调用 NPU、FSDP 看不到低精 Tensor、hook 重复注册或 optimizer
拿不到高精 master,属于启动时必须检查的集成错误。
15. 约束限制与失败规则
启动时按 policy 校验,不通过直接失败:
torch_npu.hifloat8、current quant、指定 scale quant、FP32 scale MM/GMM;TP/CP/EP/PP > 1;fallback_to_unsharded。运行阶段不得因 capability、shape 或算子不支持退回 BF16/FP32 GEMM。
完整限制矩阵:
falsefallback_to_unsharded控制复制或失败16. UT 与组件测试设计
UT 默认在 CPU/mock NPU Ops 上执行;验证 Python 契约、数据流和失败语义。真实 kernel 数值只在
A5 ST/转测验证,UT 不伪装成硬件验收。
16.1 配置与 Replacement Plan UT
format/scalingenabled=false16.2 Tensor、Quantizer 与 Ops UT
amax/format_max,零/NaN/Inf 安全处理[E]与 expert 顺序一致select(row/col)返回正确 data/scaleQuantizer 数值 UT 使用确定性小 Tensor 手算 scale 和反量化误差;MM/GMM Python UT 使用 fake
operator 记录调用参数,不把 fake 输出作为真实低精精度结论。
16.3 Dense/Grouped Autograd UT
NN/NT/TN三条 Fprop/Dgrad/Wgrad 路径和量化角色;16.4 FSDP2 Extension UT
reshard_after_forward=true/false和 prefetch phase;[E]scale;fallback_to_unsharded的复制/报错分支;comm_fusion=true组合快速失败。16.5 Optimizer、Recompute 与 DCP UT/ST
exp_avg/exp_avg_sq均为 FP32;activation_checkpoint=off/full下 loss、梯度和 kernel 路径一致,full 不重复刷新权重。16.6 A5 端到端 ST 伪代码
ST 必须经过生产 Trainer 入口。
STProbe只观测模块类型、kernel、FSDP physical tensors、scale和 refresh 次数,不在测试中手工执行 replacement、量化、all-gather 或权重刷新。
# 伪代码:接口名称表达测试职责,不要求与最终实现逐字一致。 @dataclass(frozen=True) class STCase: name: str low_precision: dict | None # None 表示 BF16 baseline reshard_after_forward: bool activation_checkpoint: str # "off" / "full" grad_accum_steps: int CASES = [ STCase("bf16", None, True, "off", 1), STCase( "mxfp8", {"enabled": True, "format": "mxfp8_e4m3", "scaling": "mx_block"}, True, "off", 1, ), STCase( "hif8_current", {"enabled": True, "format": "hif8", "scaling": "current"}, True, "off", 1, ), ] def run_st(case: STCase, checkpoint_dir: Path) -> STResult: init_process_group(backend="hccl") set_deterministic_seed(2026) # 所有 case 使用相同 seed 和 batch 顺序 cfg = load_yaml("dsv3_16b_fsdp2_st.yaml") cfg.low_precision = case.low_precision cfg.fsdp_config.reshard_after_forward = case.reshard_after_forward cfg.gradient_checkpointing.activation_checkpoint = case.activation_checkpoint cfg.gradient_accumulation_steps = case.grad_accum_steps # setup() 必须走正式顺序:HF 构建 → replacement → finalize → FSDP2 → optimizer。 trainer = Trainer() probe = STProbe(read_only=True) trainer.register_test_probe(probe) trainer.setup(cfg) assert_model_scope(trainer.model, case, probe) assert_optimizer_fp32_state(trainer.optimizer) if case.low_precision is not None: assert_no_high_precision_gemm_fallback(probe) step_records = [] for step in range(WARMUP_STEPS + MEASURE_STEPS): trainer.optimizer.zero_grad(set_to_none=True) probe.begin_optimizer_step(step) for micro_step in range(case.grad_accum_steps): batch = deterministic_batch(step, micro_step) loss = trainer.forward_backward( batch, loss_scale=1.0 / case.grad_accum_steps, ) assert torch.isfinite(loss) assert_all_gradients_finite(trainer.model) # 由生产 optimizer adapter 在成功 step 后刷新;ST 只记录事件。 updated = trainer.optimizer_step() probe.end_optimizer_step(updated=updated) if updated and case.low_precision is not None: probe.assert_one_weight_refresh_per_parameter(step) if not updated: probe.assert_no_weight_refresh(step) torch.npu.synchronize() step_records.append( collect_step_record( step=step, loss=loss, grad_norm=trainer.grad_norm, elapsed=probe.step_elapsed, peak_memory=torch.npu.max_memory_allocated(), ) ) if step == SAVE_STEP: trainer.save_dcp(checkpoint_dir) if case.low_precision is not None: assert_low_precision_kernel_contract(case, probe) assert_fsdp_physical_tensor_contract(case, probe) assert_scale_contract(case, probe) return STResult( case=case, records=step_records, checkpoint_dir=checkpoint_dir, trace=probe.export_trace(), )关键断言伪代码:
def assert_low_precision_kernel_contract(case, probe): # Dense 与 routed experts 的三种反向角色必须全部真实执行。 expected = { ("dense", "fprop"), ("dense", "dgrad"), ("dense", "wgrad"), ("grouped", "fprop"), ("grouped", "dgrad"), ("grouped", "wgrad"), } assert expected <= probe.executed_kernel_roles() assert probe.high_precision_fallback_count == 0 if case.name == "mxfp8": probe.assert_all_kernel_format("mxfp8_e4m3") probe.assert_mxfp8_scale_dtype("e8m0") elif case.name == "hif8_current": probe.assert_all_kernel_format("hif8") probe.assert_hif8_scale_dtype(torch.float32) probe.assert_hif8_format_max(input_weight=15.0, gradient=224.0) def assert_fsdp_physical_tensor_contract(case, probe): for event in probe.fsdp_gather_events: event.assert_post_matches_pre_schema() if case.name == "mxfp8": probe.assert_mxfp8_forward_direction_by_compute_kind() probe.assert_mxfp8_backward_has_row_and_col() elif case.name == "hif8_current": probe.assert_hif8_forward_backward_use_data_and_scale() probe.assert_hif8_dense_scale_equal_across_ranks() probe.assert_hif8_grouped_scale_matches_expert_order() def assert_scale_contract(case, probe): if case.name != "hif8_current": return for event in probe.hif8_quant_events: expected_max = 224.0 if event.role == "grad_output" else 15.0 expected_scale = safe_scale(event.current_amax, expected_max) torch.testing.assert_close(event.scale, expected_scale) assert event.scale.dtype == torch.float32 assert torch.isfinite(event.scale).all() and (event.scale > 0).all()DCP 续训使用新的 Trainer 实例,防止旧进程内对象掩盖恢复问题:
def run_resume_st(case, continuous_result): resumed = Trainer() resumed.setup(load_same_case_config(case)) resumed.load_dcp(continuous_result.checkpoint_dir) assert resumed.global_step == SAVE_STEP assert_fp32_master_and_moments(resumed.optimizer) assert_checkpoint_has_no_derived_low_precision_storage( continuous_result.checkpoint_dir ) # load_dcp 后由生产 adapter 从 FP32 master 重建低精权重。 assert_rebuilt_weight_storage(resumed.model, case) for step in range(SAVE_STEP, SAVE_STEP + RESUME_STEPS): batch = deterministic_batch(step, 0) resumed_record = resumed.train_step(batch) continuous_record = continuous_result.record(step) assert_resume_metrics_close(resumed_record, continuous_record)性能汇总伪代码:
bf16 = run_st(BF16_CASE, tmp_path / "bf16") mxfp8 = run_st(MXFP8_CASE, tmp_path / "mxfp8") hif8 = run_st(HIF8_CURRENT_CASE, tmp_path / "hif8") for low_precision in (mxfp8, hif8): # 丢弃 warmup,只统计相同 token 数的连续 optimizer steps。 assert low_precision.steady_tokens_per_second >= bf16.steady_tokens_per_second assert low_precision.average_step_time <= bf16.average_step_time8 卡启动形式:
torchrun --nproc_per_node=8 tests/st/low_precision/test_dsv3_16b_fsdp2.py \ --config configs/st/dsv3_16b_fsdp2.yaml \ --case hif8_current17. Q3 转测范围与验收标准
17.1 转测环境与固定输入
Q3 转测使用单机 8 卡 Ascend A5(每卡 96 GiB)、PyTorch + torch_npu,固定 DeepSeek-V3
16B 规格和同一份确定性训练数据。每次测试记录代码 commit、CANN/torch_npu/PyTorch 版本、
完整 YAML、随机种子和算子 trace。若 16B 因环境资源问题无法启动,只能用缩小模型定位,不能
以缩小模型代替最终验收。
固定配置维度:
17.2 阻塞验收矩阵
reshard_after_forward=true/false各运行 5 stepsreshard_after_forward=true/false各运行 5 stepsactivation_checkpoint=full,各格式 5 steps这里的“约定容差”必须在首轮 A5 BF16、MXFP8、HiF8 固定 seed 数据采集后,由算法负责人给出
具体数值并写回转测用例;在阈值冻结前,Q3-09 不得仅凭“曲线看起来接近”判定通过。
17.3 数值与精度专项
17.3.1 HiF8 专项验收范围
Q3 HiF8 专项验收仅覆盖
hif8 + current,不覆盖 delayed scaling。转测不需要构造、检查或恢复 amax history、pending amax、scale update interval 等 delayed 状态。
[E_local,...]data、[E_local]scale 的 gatherHiF8 delayed scaling 配置必须在启动阶段报“不支持”,而不是被忽略、映射为 current 或回退到
高精计算。delayed scaling 的实现与验收统一放到第 18 章后续演进。
17.3.2 分层数值验收
端到端 loss 不能单独证明 Fprop/Dgrad/Wgrad 正确,必须分层验收:
abs、MSE、NRMSE、cosine 和 nonfinite count。
Inf、持续爆炸或从第一个 step 起完全不更新。
optimizer step。
Q3 首轮允许“采集阈值”与功能转测并行,但发布前必须冻结上述指标的数值门槛。文档不预设
未经 A5 实测的统一误差阈值,避免用任意阈值掩盖格式和算子差异。
17.4 性能与显存专项
性能是 Q3 阻塞验收项。比较时必须固定模型、训练数据、sequence length、micro/global batch、
梯度累积、FSDP、reshard 和 recompute 配置;不得通过缩小 batch、减少有效 token 或改变并行
策略得到表面加速。关闭 profiling 和一次性初始化影响,预热后至少采集 20 个连续 optimizer
step,以相同统计方法计算:
MXFP8 和 HiF8 current 必须分别满足上述条件,不能用一种格式的收益抵消另一种格式的劣化。
同时必须产出:
reshard_after_forward和 recompute 开关的差异;任一低精格式吞吐低于 BF16 或平均 step time 高于 BF16,均判定 Q3 性能验收失败。峰值显存
本期要求记录并解释,暂不设置相对 BF16 的阻塞阈值;若发现错误 fallback、重复量化、重复
all-gather 或缓存未释放,则同时按功能问题阻塞。
17.5 转测交付物
18. 后续演进
HiF8 delayed scaling 后续单独设计和交付。届时需要新增角色级 amax history、跨 rank amax
归并、optimizer-step 更新时序和 DCP scaling-state 保存恢复;本期不提前引入这些状态、配置或
生命周期接口。