已合并
fix(moe): 补齐 MoeFinalizeRouting / MoeReRouting 的结构化输入契约 #251
Deng Pan创建于 16 天前
fix(moe): 补齐 MoeFinalizeRouting / MoeReRouting 的结构化输入契约 #251
已合并
共 3 个文件变更+276-1
| @@ -311,4 +311,73 @@ if __name__ == "__main__": | |||
| 311 | golden_torch = moe_finalize_routing(**inputs["torch"]) | 311 | golden_torch = moe_finalize_routing(**inputs["torch"]) |
| 312 | print(f"Torch golden output shape: {golden_torch.shape}") | 312 | print(f"Torch golden output shape: {golden_torch.shape}") |
| 313 | print(f"Torch golden output dtype: {golden_torch.dtype}") | 313 | print(f"Torch golden output dtype: {golden_torch.dtype}") |
| 314 | - print(f"Torch golden output sample: {golden_torch[0, :5]}") | 314 | + print(f"Torch golden output sample: {golden_torch[0, :5]}") |
| 315 | + | ||
| 316 | +def get_input( | ||
| 317 | + expanded_permuted_rows: torch.Tensor, | ||
| 318 | + expanded_src_to_dst_row: torch.Tensor, | ||
| 319 | + skip1: Optional[torch.Tensor] = None, | ||
| 320 | + skip2: Optional[torch.Tensor] = None, | ||
| 321 | + bias: Optional[torch.Tensor] = None, | ||
| 322 | + scales: Optional[torch.Tensor] = None, | ||
| 323 | + expert_for_source_row: Optional[torch.Tensor] = None, | ||
| 324 | + drop_pad_mode: int = 0, | ||
| 325 | + **kwargs, | ||
| 326 | +): | ||
| 327 | + """把 expanded_src_to_dst_row 重建为满足单射契约的行映射。 | ||
| 328 | + | ||
| 329 | + 该输入是 MoeInitRouting 产出的"源行 -> 展开缓冲区行"映射:每个源行 (token, k) | ||
| 330 | + 占据展开缓冲区中**互不相同**的一行。drop_less (mode 0/2) 下它是 [0, NK) 的完整 | ||
| 331 | + 置换;drop_pad (mode 1/3) 下是到 [0, E*C) 的单射,容量溢出的源行取 -1 表示丢弃。 | ||
| 332 | + | ||
| 333 | + 而单区间 value_range 表达不了"互异",通用生成器按 randint 独立有放回采样,导致 | ||
| 334 | + 部分展开行被读多次、另一部分从未被读。golden 与真实 MoeFinalizeRoutingV2 都做 | ||
| 335 | + gather,重复索引下数值仍然可比(现网用例即如此通过),但: | ||
| 336 | + 1. 访存模式与真实场景不符——真实场景每行恰好读一次,重复采样把它变成随机重复 | ||
| 337 | + 读,L2 命中率虚高,perf 数据不具代表性; | ||
| 338 | + 2. 任何采用 scatter 方向实现的候选(遍历展开行写回目的行,与 gather 等价当且 | ||
| 339 | + 仅当映射单射)都会与 golden 分叉,而契约本身是站在候选一边的。 | ||
| 340 | + | ||
| 341 | + 这里按 case 的实际形状推导目的空间大小 num_dst = expanded_permuted_rows 展平成 | ||
| 342 | + (num_dst, H) 后的行数(drop_less 下等于 NK,drop_pad 下等于 E*C),为非 -1 的位置 | ||
| 343 | + 分配 [0, num_dst) 的互异值。**-1 的位置原样保留**,以维持 drop_pad 用例既有的 | ||
| 344 | + 丢弃覆盖与随种子可复现的行为。 | ||
| 345 | + | ||
| 346 | + 注:drop_pad 的完整契约还要求目的行落在该源行所属专家的容量块 [e*C, (e+1)*C) 内 | ||
| 347 | + (e 取自 expert_for_source_row)。golden 不校验这一点,且按专家分配会让丢弃率降到 | ||
| 348 | + 近乎 0(当前用例 E*C 远大于 NK),反而削弱 -1 路径覆盖,故此处只做单射重建;如需 | ||
| 349 | + 专家对齐应由用例设计一并调整容量。 | ||
| 350 | + | ||
| 351 | + kernel_eval 用输入名 + attrs 作为关键字调用本函数,并用返回值(按 golden 签名的 | ||
| 352 | + Tensor 顺序)同时替换 golden 与候选的输入,故比较公平。 | ||
| 353 | + | ||
| 354 | + Returns: | ||
| 355 | + [expanded_permuted_rows, expanded_src_to_dst_row, skip1, skip2, bias, scales, | ||
| 356 | + expert_for_source_row],顺序与 moe_finalize_routing 签名一致。 | ||
| 357 | + """ | ||
| 358 | + unchanged = [expanded_permuted_rows, expanded_src_to_dst_row, skip1, skip2, | ||
| 359 | + bias, scales, expert_for_source_row] | ||
| 360 | + esdr = expanded_src_to_dst_row | ||
| 361 | + if not isinstance(esdr, torch.Tensor) or esdr.numel() == 0: | ||
| 362 | + return unchanged | ||
| 363 | + | ||
| 364 | + H = int(expanded_permuted_rows.shape[-1]) | ||
| 365 | + num_dst = expanded_permuted_rows.numel() // H | ||
| 366 | + | ||
| 367 | + flat = esdr.reshape(-1) | ||
| 368 | + keep = flat != -1 | ||
| 369 | + n_keep = int(keep.sum()) | ||
| 370 | + if n_keep > num_dst: | ||
| 371 | + # 抽屉原理:单射不存在(当前用例集不会走到这里),保持原样 | ||
| 372 | + return unchanged | ||
| 373 | + | ||
| 374 | + g = torch.Generator().manual_seed(0) # 固定种子:跨 eval 运行必须可复现 | ||
| 375 | + # argsort(rand) 即随机置换;取前 n_keep 项得到 [0, num_dst) 内互异的目的行 | ||
| 376 | + dst = torch.rand(num_dst, generator=g).argsort()[:n_keep] | ||
| 377 | + | ||
| 378 | + new_flat = flat.clone() | ||
| 379 | + new_flat[keep] = dst.to(dtype=flat.dtype, device=flat.device) | ||
| 380 | + new_esdr = new_flat.reshape(esdr.shape) | ||
| 381 | + | ||
| 382 | + return [expanded_permuted_rows, new_esdr, skip1, skip2, bias, scales, | ||
| 383 | + expert_for_source_row] | ||
| @@ -160,6 +160,18 @@ def get_input( | |||
| 160 | base_value = A // total_cells | 160 | base_value = A // total_cells |
| 161 | remainder = A % total_cells | 161 | remainder = A % total_cells |
| 162 | 162 | ||
| 163 | + # proto.yaml 声明 expert_token_num_per_rank 的元素**必须大于 0**。A < N*E 时 | ||
| 164 | + # base_value 为 0,除最后一格外全是 0 —— 静默产出违反契约的输入,候选 kernel | ||
| 165 | + # 在"某张卡的某个专家分到 0 个 token"上的行为是未定义的,且失败会表现为莫名的 | ||
| 166 | + # 精度不符而非配置错误。当前用例集最小 base_value = 4,不会走到这里;这是为 | ||
| 167 | + # 后续新增小 A 用例设的护栏,宁可显式报错也不静默失效。 | ||
| 168 | + if base_value < 1: | ||
| 169 | + raise ValueError( | ||
| 170 | + f"moe_re_routing get_input: tokens 数 A={A} 少于 expert_token_num_per_rank " | ||
| 171 | + f"的格子数 N*E={N}*{E}={total_cells},无法让每格都 >0(proto 要求元素必须" | ||
| 172 | + f"大于 0)。请调整该用例的 shape 使 A >= N*E。" | ||
| 173 | + ) | ||
| 174 | + | ||
| 163 | # 生成新的 expert_token_num_per_rank | 175 | # 生成新的 expert_token_num_per_rank |
| 164 | if isinstance(expert_token_num_per_rank, torch.Tensor): | 176 | if isinstance(expert_token_num_per_rank, torch.Tensor): |
| 165 | new_expert_token_num = torch.full((N, E), base_value, dtype=expert_token_num_per_rank.dtype) | 177 | new_expert_token_num = torch.full((N, E), base_value, dtype=expert_token_num_per_rank.dtype) |
| @@ -0,0 +1,194 @@ | |||
| 1 | +#!/usr/bin/python3 | ||
| 2 | +# coding=utf-8 | ||
| 3 | + | ||
| 4 | +# ---------------------------------------------------------------------------------------------------------- | ||
| 5 | +# Copyright (c) 2026 Huawei Technologies Co., Ltd. | ||
| 6 | +# This program is free software; you can redistribute it and/or modify it under the terms and conditions of | ||
| 7 | +# CANN Open Software License Agreement Version 2.0 (the "License"). | ||
| 8 | +# Please refer to the License for details. You may not use this file except in compliance with the License. | ||
| 9 | +# THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED, | ||
| 10 | +# INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE. | ||
| 11 | +# See LICENSE in the root of the software repository for the full text of the License. | ||
| 12 | +# ---------------------------------------------------------------------------------------------------------- | ||
| 13 | + | ||
| 14 | +"""MoE 类算子结构化输入契约的 get_input 测试 | ||
| 15 | + | ||
| 16 | +测试对象: | ||
| 17 | +1. tasks/level3/moe_finalize_routing/golden.py::get_input | ||
| 18 | + —— expanded_src_to_dst_row 是"源行 -> 展开缓冲区行"的单射映射(drop_less 下是 | ||
| 19 | + 完整置换),单区间 value_range 表达不了"互异",通用生成器有放回采样出大量重复。 | ||
| 20 | +2. tasks/level3/moe_re_routing/golden.py::get_input | ||
| 21 | + —— proto 要求 expert_token_num_per_rank 元素必须大于 0,A < N*E 时原实现静默产出 0。 | ||
| 22 | +""" | ||
| 23 | + | ||
| 24 | +import importlib.util | ||
| 25 | + | ||
| 26 | +import pytest | ||
| 27 | +import torch | ||
| 28 | + | ||
| 29 | +from kernel_eval.config import get_project_root | ||
| 30 | + | ||
| 31 | + | ||
| 32 | +def _load(op_rel): | ||
| 33 | + path = get_project_root() / "tasks" / op_rel / "golden.py" | ||
| 34 | + spec = importlib.util.spec_from_file_location(f"_golden_{op_rel.replace('/', '_')}", path) | ||
| 35 | + mod = importlib.util.module_from_spec(spec) | ||
| 36 | + spec.loader.exec_module(mod) | ||
| 37 | + return mod | ||
| 38 | + | ||
| 39 | + | ||
| 40 | + | ||
| 41 | +def finalize_routing(): | ||
| 42 | + return _load("level3/moe_finalize_routing") | ||
| 43 | + | ||
| 44 | + | ||
| 45 | + | ||
| 46 | +def re_routing(): | ||
| 47 | + return _load("level3/moe_re_routing") | ||
| 48 | + | ||
| 49 | + | ||
| 50 | +# --------------------------------------------------------------------------- | ||
| 51 | +# MoeFinalizeRouting:expanded_src_to_dst_row 单射契约 | ||
| 52 | +# --------------------------------------------------------------------------- | ||
| 53 | + | ||
| 54 | + | ||
| 55 | +def _drop_less_case(num_rows=16, k=4, hidden=8, seed=0): | ||
| 56 | + """drop_less(mode 0/2):num_dst == NK,映射应为 [0, NK) 的完整置换""" | ||
| 57 | + g = torch.Generator().manual_seed(seed) | ||
| 58 | + nk = num_rows * k | ||
| 59 | + epr = torch.rand(nk, hidden, generator=g) | ||
| 60 | + esdr = torch.randint(0, nk, (nk,), generator=g, dtype=torch.int32) | ||
| 61 | + scales = torch.rand(num_rows, k, generator=g) | ||
| 62 | + skip1 = torch.rand(num_rows, hidden, generator=g) | ||
| 63 | + return epr, esdr, scales, skip1 | ||
| 64 | + | ||
| 65 | + | ||
| 66 | +def _drop_pad_case(num_rows=32, k=1, experts=8, capacity=10, hidden=8, seed=0): | ||
| 67 | + """drop_pad(mode 1/3):epr 为 (E, C, H),num_dst = E*C > NK,-1 表示丢弃""" | ||
| 68 | + g = torch.Generator().manual_seed(seed) | ||
| 69 | + nk = num_rows * k | ||
| 70 | + epr = torch.rand(experts, capacity, hidden, generator=g) | ||
| 71 | + esdr = torch.randint(-1, experts * capacity, (nk,), generator=g, dtype=torch.int32) | ||
| 72 | + scales = torch.rand(num_rows, k, generator=g) | ||
| 73 | + skip1 = torch.rand(num_rows, hidden, generator=g) | ||
| 74 | + return epr, esdr, scales, skip1 | ||
| 75 | + | ||
| 76 | + | ||
| 77 | +class TestFinalizeRoutingInjectivity: | ||
| 78 | + def test_drop_less_becomes_full_permutation(self, finalize_routing): | ||
| 79 | + epr, esdr, scales, skip1 = _drop_less_case() | ||
| 80 | + nk = esdr.numel() | ||
| 81 | + assert esdr.unique().numel() < nk, "构造失败:原始映射本应有重复" | ||
| 82 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=0) | ||
| 83 | + new = out[1] | ||
| 84 | + assert sorted(new.tolist()) == list(range(nk)) | ||
| 85 | + | ||
| 86 | + def test_drop_pad_non_sentinel_entries_distinct(self, finalize_routing): | ||
| 87 | + epr, esdr, scales, skip1 = _drop_pad_case() | ||
| 88 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=1) | ||
| 89 | + new = out[1] | ||
| 90 | + keep = new != -1 | ||
| 91 | + assert int(keep.sum()) > 0 | ||
| 92 | + assert new[keep].unique().numel() == int(keep.sum()) | ||
| 93 | + | ||
| 94 | + def test_drop_pad_sentinel_positions_preserved(self, finalize_routing): | ||
| 95 | + """-1 的位置原样保留,维持既有的丢弃覆盖""" | ||
| 96 | + epr, esdr, scales, skip1 = _drop_pad_case() | ||
| 97 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=1) | ||
| 98 | + assert torch.equal(esdr == -1, out[1] == -1) | ||
| 99 | + | ||
| 100 | + def test_indices_within_destination_space(self, finalize_routing): | ||
| 101 | + epr, esdr, scales, skip1 = _drop_pad_case() | ||
| 102 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=1) | ||
| 103 | + new = out[1] | ||
| 104 | + keep = new != -1 | ||
| 105 | + num_dst = epr.numel() // epr.shape[-1] | ||
| 106 | + assert int(new[keep].min()) >= 0 | ||
| 107 | + assert int(new[keep].max()) < num_dst | ||
| 108 | + | ||
| 109 | + def test_shape_and_dtype_unchanged(self, finalize_routing): | ||
| 110 | + epr, esdr, scales, skip1 = _drop_less_case() | ||
| 111 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=0) | ||
| 112 | + assert out[1].shape == esdr.shape | ||
| 113 | + assert out[1].dtype == esdr.dtype | ||
| 114 | + | ||
| 115 | + def test_reproducible_across_calls(self, finalize_routing): | ||
| 116 | + epr, esdr, scales, skip1 = _drop_less_case() | ||
| 117 | + a = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=0)[1] | ||
| 118 | + b = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=0)[1] | ||
| 119 | + assert torch.equal(a, b) | ||
| 120 | + | ||
| 121 | + def test_other_inputs_passed_through(self, finalize_routing): | ||
| 122 | + epr, esdr, scales, skip1 = _drop_less_case() | ||
| 123 | + bias = torch.rand(4, epr.shape[-1]) | ||
| 124 | + efsr = torch.randint(0, 4, (16, 4), dtype=torch.int32) | ||
| 125 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, skip2=None, bias=bias, | ||
| 126 | + scales=scales, expert_for_source_row=efsr, | ||
| 127 | + drop_pad_mode=0) | ||
| 128 | + assert len(out) == 7 | ||
| 129 | + assert out[0] is epr and out[2] is skip1 and out[3] is None | ||
| 130 | + assert out[4] is bias and out[5] is scales and out[6] is efsr | ||
| 131 | + | ||
| 132 | + def test_extra_kwargs_tolerated(self, finalize_routing): | ||
| 133 | + epr, esdr, scales, skip1 = _drop_less_case() | ||
| 134 | + finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, | ||
| 135 | + drop_pad_mode=0, skip2_exist=True) | ||
| 136 | + | ||
| 137 | + def test_impossible_injection_passthrough(self, finalize_routing): | ||
| 138 | + """抽屉原理:源行多于目的行时保持原样而非报错""" | ||
| 139 | + g = torch.Generator().manual_seed(0) | ||
| 140 | + epr = torch.rand(4, 8, generator=g) # num_dst = 4 | ||
| 141 | + esdr = torch.randint(0, 4, (16,), generator=g, dtype=torch.int32) | ||
| 142 | + out = finalize_routing.get_input(epr, esdr, drop_pad_mode=0) | ||
| 143 | + assert out[1] is esdr | ||
| 144 | + | ||
| 145 | + def test_gather_semantics_unchanged(self, finalize_routing): | ||
| 146 | + """重建后 golden 仍可正常执行,输出规格不变""" | ||
| 147 | + epr, esdr, scales, skip1 = _drop_less_case() | ||
| 148 | + out = finalize_routing.get_input(epr, esdr, skip1=skip1, scales=scales, drop_pad_mode=0) | ||
| 149 | + y = finalize_routing.moe_finalize_routing(out[0], out[1], skip1=out[2], scales=out[5], | ||
| 150 | + drop_pad_mode=0) | ||
| 151 | + assert y.shape == skip1.shape | ||
| 152 | + | ||
| 153 | + | ||
| 154 | +# --------------------------------------------------------------------------- | ||
| 155 | +# MoeReRouting:expert_token_num_per_rank 元素必须 > 0 | ||
| 156 | +# --------------------------------------------------------------------------- | ||
| 157 | + | ||
| 158 | + | ||
| 159 | +class TestReRoutingGuard: | ||
| 160 | + def test_feasible_case_sums_to_token_count(self, re_routing): | ||
| 161 | + tokens = torch.rand(1024, 16) | ||
| 162 | + etn = torch.zeros(8, 8, dtype=torch.int32) | ||
| 163 | + out = re_routing.get_input(tokens, etn) | ||
| 164 | + assert int(out[1].sum()) == tokens.shape[0] | ||
| 165 | + assert int(out[1].min()) > 0 | ||
| 166 | + | ||
| 167 | + def test_remainder_still_absorbed(self, re_routing): | ||
| 168 | + """A 不整除 N*E 时总和仍须等于 A(沿用原有的"余数并入最后一格")""" | ||
| 169 | + tokens = torch.rand(1009, 16) | ||
| 170 | + etn = torch.zeros(8, 8, dtype=torch.int32) | ||
| 171 | + out = re_routing.get_input(tokens, etn) | ||
| 172 | + assert int(out[1].sum()) == 1009 | ||
| 173 | + assert int(out[1].min()) > 0 | ||
| 174 | + | ||
| 175 | + def test_infeasible_case_raises_instead_of_emitting_zeros(self, re_routing): | ||
| 176 | + """A < N*E:无法让每格都 >0,应显式报错而非静默产出违反契约的 0""" | ||
| 177 | + tokens = torch.rand(32, 16) # A=32 < N*E=64 | ||
| 178 | + etn = torch.zeros(8, 8, dtype=torch.int32) | ||
| 179 | + with pytest.raises(ValueError, match="N\\*E"): | ||
| 180 | + re_routing.get_input(tokens, etn) | ||
| 181 | + | ||
| 182 | + def test_boundary_a_equals_cells(self, re_routing): | ||
| 183 | + """A == N*E 恰好每格 1 个,属可行边界""" | ||
| 184 | + tokens = torch.rand(64, 16) | ||
| 185 | + etn = torch.zeros(8, 8, dtype=torch.int32) | ||
| 186 | + out = re_routing.get_input(tokens, etn) | ||
| 187 | + assert int(out[1].min()) == 1 and int(out[1].sum()) == 64 | ||
| 188 | + | ||
| 189 | + def test_per_token_scales_passed_through(self, re_routing): | ||
| 190 | + tokens = torch.rand(1024, 16) | ||
| 191 | + etn = torch.zeros(8, 8, dtype=torch.int32) | ||
| 192 | + scales = torch.rand(1024) | ||
| 193 | + out = re_routing.get_input(tokens, etn, per_token_scales=scales) | ||
| 194 | + assert out[0] is tokens and out[2] is scales | ||