已合并
fix(moe): 补齐 MoeFinalizeRouting / MoeReRouting 的结构化输入契约 #251
fix(moe): 补齐 MoeFinalizeRouting / MoeReRouting 的结构化输入契约 #251
已合并
Deng Pan创建于 16 天前
3 个文件变更+276-1
Mtasks/level3/moe_finalize_routing/golden.py+70-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]
Mtasks/level3/moe_re_routing/golden.py+12-0
@@ -160,6 +160,18 @@ def get_input(
160 base_value = A // total_cells160 base_value = A // total_cells
161 remainder = A % total_cells161 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_rank175 # 生成新的 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)
Atests/ut/test_moe_get_input.py+194-0
@@ -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+@pytest.fixture(scope="module")
41+def finalize_routing():
42+ return _load("level3/moe_finalize_routing")
43+ 
44+ 
45+@pytest.fixture(scope="module")
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