已合并
[A5]: Adding support for TMOV DN to ZZ #1196
omarzohir创建于 6月26日
[A5]: Adding support for TMOV DN to ZZ #1196
已合并
omarzohir创建于 6月26日
共 7 个文件变更+518-168
@@ -0,0 +1,159 @@
1+# TQUANT DN — Axis-0 Grouped Quantization and DN→ZZ
2+ 
3+## Tile Operation Diagram
4+ 
5+![TQUANT tile operation](../figures/isa/TQUANT.svg)
6+ 
7+## Introduction
8+ 
9+**DN** (Down-column Normal) denotes MX quantization with groups of 32 along **axis 0**
10+(rows), as opposed to the default **ND** (Normal) which groups along **axis 1** (columns).
11+Both modes produce RowMajor tiles — "DN" refers only to the *grouping direction*, not
12+the storage layout. DN is used in FlashAttention where the softmax output (P matrix)
13+has its natural grouping along the M (row) dimension.
14+ 
15+After DN quantization, the FP8 data is converted to NZ via the stock `TMOV(ND→NZ)` and
16+the E8M0 exponents are converted to ZZ via the new `TMOV<grp_axis=0>(DN→ZZ)`.
17+ 
18+## C++ Intrinsic
19+ 
20+The primary interface is the `<grp_axis, mx_alg>` template:
21+ 
22+```cpp
23+template <int grp_axis, auto mx_alg, typename TileDataOut = void, typename TileDataSrc = void,
24+ typename TileDataExp = void, typename TileDataMax = void, typename TileDataScaling = void,
25+ typename... WaitEvents>
26+PTO_INST RecordEvent TQuant(TileDataOut &dst, TileDataSrc &src, TileDataExp *exp, TileDataMax *max,
27+ TileDataScaling *scaling, WaitEvents &... events);
28+```
29+ 
30+### Parameters
31+ 
32+| Parameter | Description |
33+|-----------|-------------|
34+| `grp_axis` | **0** = DN (groups on axis 0 / rows); **1** = ND (groups on axis 1 / columns, default) |
35+| `mx_alg` | Combined destination-format + scale-algorithm tag (`MxQuantAlg` enum) |
36+| `dst` | Output FP8/FP4 tile (RowMajor, same shape as `src`) |
37+| `src` | Input fp32/bf16/fp16 tile (RowMajor `M×N`) |
38+| `exp` | Output E8M0 exponent tile: shape `M̂×N` for DN, `M×Γ` for ND |
39+| `max` | Scratch per-group abs-max tile |
40+| `scaling` | Scratch per-group scaling tile |
41+ 
42+### `MxQuantAlg` values
43+ 
44+```cpp
45+enum class MxQuantAlg {
46+ OcpMxFp8E4M3 = 0, // MXFP8 E4M3 + OCP scale
47+ NvMxFp8E4M3 = 1, // MXFP8 E4M3 + NV scale
48+ OcpMxFp4E2M1 = 2, // MXFP4 E2M1 + OCP scale
49+ NvMxFp4E2M1 = 3, // MXFP4 E2M1 + NV scale
50+};
51+```
52+ 
53+> **Backward compatibility:** the old `TQUANT<QuantType::MXFP8, ...>` interface is
54+> retained unchanged. The `<grp_axis, mx_alg>` form is the preferred interface going
55+> forward; nothing is removed.
56+ 
57+## DN Output Shapes
58+ 
59+For a source tile `M×N` with `M̂ = M/32`, `Γ = N/32`:
60+ 
61+| Output | ND (`grp_axis=1`) | DN (`grp_axis=0`) |
62+|--------|-------------------|-------------------|
63+| FP8/FP4 data | `M×N` RowMajor | `M×N` RowMajor (identical) |
64+| E8M0 exponent | `M×Γ` | `M̂×N` |
65+| Max / Scaling | `M×Γ` | `M̂×N` |
66+ 
67+The **data tile is identical** between ND and DN (same `(r,c)` addresses); only the
68+exponent/max/scaling tile shapes differ. Therefore `TMOV(ND→NZ)` on the data is
69+reused unchanged. Only the exponent needs a new transform: **DN→ZZ**.
70+ 
71+## Cube Consumption Contract
72+ 
73+Verified from A5 sim logs (`LOAD_2Dv2` + `LOAD_MX_2Dv2` + `MMAD_MX`):
74+ 
75+```
76+FP8 data → L0A/L0B as NZ fractal (LOAD_2Dv2 Dtype:B8)
77+E8M0 scale → L0AMX/L0BMX as ZZ fractal (LOAD_MX_2Dv2 Dtype:B16)
78+MMAD_MX pairs them by fractal byte position.
79+```
80+ 
81+The cube always wants data in NZ and scale in ZZ, regardless of the quantization group
82+axis. The only difference for DN-quantized operands is the exponent tile shape
83+(`M̂×N` vs `M×Γ`) and the transform applied (`DN→ZZ` vs `ND→ZZ`).
84+ 
85+## DN→ZZ Transformation
86+ 
87+### Mathematical Proof
88+ 
89+For DN exponent tile `E_DN[hat_r][c]` of shape `M̂×N` (RowMajor, flat `hat_r·N + c`):
90+ 
91+**Theorem (DN→ZZ = transpose ⊕ ND→ZZ):**
92+ 
93+$$E_{ZZ}[c_b, p, q, \delta] = E_{DN}^T[16c_b + q][2p + \delta] = E_{DN}[2p + \delta][16c_b + q]$$
94+ 
95+with `c_b ∈ [0, N/16)`, `p ∈ [0, M̂/2)`, `q ∈ [0,16)`, `δ ∈ {0,1}`.
96+ 
97+**Corollary (direct source index):**
98+ 
99+$$\text{src\_idx}(c_b, p, q, \delta) = (2p + \delta) \cdot N + 16c_b + q$$
100+ 
101+**Corollary (no gather needed):** For fixed `(c_b, p)`, the 32 bytes of the ZZ box are
102+sourced from two contiguous 16-byte runs: `E_DN[2p][16c_b:16c_b+16]` and
103+`E_DN[2p+1][16c_b:16c_b+16]`. Zipping them via `vintlv` yields the `qδ`-interleaved
104+order the ZZ fractal requires. Hence DN→ZZ is cheaper than ND→ZZ (contiguous loads, no
105+`vgather2`/`BLK`/`E2B`).
106+ 
107+### Alignment Constraints
108+ 
109+- `N mod 16 = 0` (always satisfied since `N mod 32 = 0`).
110+- `M̂ mod 2 = 0`, i.e. **`M mod 64 = 0`** (for δ-pairing). Stricter than ND→ZZ's `M mod 16 = 0`.
111+- `M = 32` (`M̂ = 1`): degenerate identity (no pairs).
112+ 
113+### Relationship to vshls+vor
114+ 
115+The FA fused softmax macro's `vshls+vor` byte-pack is, *at `M̂=4, N=64` only*,
116+mathematically identical to the transpose step of this recipe. It cannot generalize
117+(requires `M̂≤4` to fit a B32 word, and `N≤64` for single-VL). `TMovDnTo2Zz` is the
118+general replacement.
119+ 
120+## TMOV Interface
121+ 
122+### DN→ZZ (new)
123+ 
124+```cpp
125+template <int grp_axis, typename DstTileData, typename SrcTileData, typename TmpTileData, typename... WaitEvents>
126+PTO_INST RecordEvent TMOV(DstTileData &dst, SrcTileData &src, TmpTileData &tmp, WaitEvents &... events);
127+```
128+ 
129+`TMOV<0>(zzTile, e8DnTile, tmpTile)` selects `TMovDnTo2Zz`. The stock
130+`TMOV(zzTile, e8Tile, tmpTile)` (without `<grp_axis>`) remains ND→ZZ (`grp_axis` defaults to 1).
131+ 
132+### ND→NZ (data, unchanged)
133+ 
134+```cpp
135+TMOV(fp8NZTile, fp8Tile); // stock 2-arg ND→NZ; correct for DN data (RowMajor, identical addresses)
136+```
137+ 
138+## Pipeline (full DN flow)
139+ 
140+```
141+src[M×N] (fp32)
142+ ──TQuant<0, MxQuantAlg::OcpMxFp8E4M3>──▶ fp8[M×N] + e8[M̂×N] (DN exponent)
143+ ──TMOV(ND→NZ)──────────────────────────▶ fp8NZ
144+ ──TMOV<0>(DN→ZZ)───────────────────────▶ e8ZZ
145+ ──feed to cube MMAD_MX──────────────────▶ C[M×N]
146+```
147+ 
148+## Examples
149+ 
150+```cpp
151+// DN quantize (groups on axis 0)
152+TQuant<0, MxQuantAlg::OcpMxFp8E4M3>(fp8Tile, srcTile, &e8DnTile, &maxTile, &scalingTile);
153+// Data ND→NZ (stock)
154+TMOV(fp8NZTile, fp8Tile);
155+// Exponent DN→ZZ (new)
156+TMOV<0>(e8ZzTile, e8DnTile, tmpTile);
157+```
158+ 
159+See `tests/npu/a5/src/st/testcase/tquant_dn/` for a complete ST example (Stages 1–3).
@@ -1275,6 +1275,18 @@ PTO_INST RecordEvent TMOV(DstTileData &dst, SrcTileData &src, TmpTileData &tmp,
1275 return {};1275 return {};
1276}1276}
1277 1277 
1278+// grp_axis-tagged X->ZZ overload (3-arg form). grp_axis=0 selects DN->ZZ on an
1279+// axis-0-grouped (M̂×N) exponent source; grp_axis=1 (default) keeps stock ND->ZZ.
1280+// Only the ZZ transform is parameterised; other TMOV overloads are unchanged.
1281+template <int grp_axis, typename DstTileData, typename SrcTileData, typename TmpTileData, typename... WaitEvents,
1282+ std::enable_if_t<is_tile_data_v<TmpTileData>, int> = 0>
1283+PTO_INST RecordEvent TMOV(DstTileData &dst, SrcTileData &src, TmpTileData &tmp, WaitEvents &...events)
1284+{
1285+ TSYNC(events...);
1286+ TMOV_IMPL<grp_axis, DstTileData, SrcTileData, TmpTileData>(dst, src, tmp);
1287+ return {};
1288+}
1289+ 
1278template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, typename... WaitEvents>1290template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, typename... WaitEvents>
1279PTO_INST RecordEvent TMOV(DstTileData &dst, SrcTileData &src, WaitEvents &...events)1291PTO_INST RecordEvent TMOV(DstTileData &dst, SrcTileData &src, WaitEvents &...events)
1280{1292{
@@ -377,6 +377,86 @@ __tf__ PTO_INTERNAL void TMovNdTo2Zz(typename DstTileData::TileDType __out__ dst
377 }377 }
378}378}
379 379 
380+// DN->ZZ gather. Source is the RowMajor M̂xN exponent tile from TQuant DN (groups on
381+// axis 0). Target ZZ box is [16,2]: 16 along N (columns), 2 along M̂ (row-group pairs):
382+// E_ZZ[cb,p,q,δ] = E_DN[2p+δ][16cb+q]
383+// The two source rows for a (cb,p) box — E_DN[2p][16cb+·] and E_DN[2p+1][16cb+·] — are
384+// each 16 contiguous bytes; vintlv zips them into the q0δ0,q0δ1,... order the ZZ box
385+// needs. So each box is two contiguous 16-B loads + vintlv + one 32-B store: no index
386+// gather, no BLK/E2B (cheaper than ND->ZZ, whose 16 intra-box elements are strided).
387+// Iterating over col blocks (cb) along a row advances the source pointer by 16 (cols
388+// are contiguous); advancing to the next row-group pair advances it by N. tmp is unused
389+// (kept in the signature to match the ND->ZZ TMOV interface).
390+template <typename DstTileData, typename SrcTileData>
391+PTO_INTERNAL void GenerateB8IndicesDN2ZZToUB(__ubuf__ uint8_t *dst, __ubuf__ uint8_t *src, __ubuf__ uint8_t *tmp,
392+ unsigned hatM, unsigned colsN)
393+{
394+ (void)tmp;
395+ const uint16_t N = (uint16_t)colsN;
396+ const uint16_t colBlockCount = N / 16;
397+ const uint16_t numPairs = (uint16_t)hatM / 2;
398+ // A ZZ box is [16,2] = 32 B = 16 B16 pairs. Store at B16 granularity with a full
399+ // predicate (PAT_ALL) and POST_UPDATE advancing by boxB16 (32 B) per box -- this is
400+ // the idiom GenerateB8IndicesZZToUB uses (store-dist-modes.md §NORM_BX: count is the
401+ // POST_UPDATE stride, predicate is the write mask; a by-value __ubuf__* local still
402+ // advances because the intrinsic takes it by reference internally).
403+ uint32_t boxB16 = 16;
404+ __ubuf__ uint16_t *dstB16 = (__ubuf__ uint16_t *)dst;
405+ // PAT_VL16 masks the lower 16 B16 lanes = exactly one 32-B ZZ box, so each store
406+ // writes only its own box (no neighbour overwrite) and POST_UPDATE advances dstB16
407+ // by boxB16 (32 B) per iteration. (PAT_ALL would write a full 256-B VL per box,
408+ // stomping the next boxes' regions.)
409+ MaskReg preg_box = pset_b16(PAT_VL16);
410+ 
411+ __ubuf__ uint8_t *rowBase = src; // E_DN[0][0]
412+ for (uint16_t cb = 0; cb < colBlockCount; ++cb) {
413+ const uint16_t off = cb * 16;
414+ for (uint16_t p = 0; p < numPairs; ++p) {
415+ // Load the two DN source rows (16 B8 each) via unaligned loads (rows are
416+ // 16-B aligned, not always 32-B). These reads were verified correct in sim.
417+ RegTensor<uint8_t> vRow0, vRow1;
418+ vector_u8 vZ0, vZ1;
419+ vector_align ureg0, ureg1;
420+ __ubuf__ uint8_t *row0 = rowBase + (uint32_t)p * 2 * N + off;
421+ __ubuf__ uint8_t *row1 = row0 + N;
422+ vldas(ureg0, row0);
423+ vldus(vRow0, ureg0, row0);
424+ vldas(ureg1, row1);
425+ vldus(vRow1, ureg1, row1);
426+ vintlv(vZ0, vZ1, (vector_u8 &)vRow0, (vector_u8 &)vRow1);
427+ // Store the first 16 B16 (32 B) of the zipped stream as B16 with a FULL
428+ // predicate + POST_UPDATE advancing dstB16 by boxB16 (32 B) per box. This
429+ // is the GenerateB8IndicesZZToUB idiom (store-dist-modes.md §NORM_BX).
430+ vsts((vector_u16 &)vZ0, dstB16, boxB16, NORM_B16, preg_box, POST_UPDATE);
431+ }
432+ }
433+}
434+ 
435+// DN->ZZ: axis-0-grouped (DN) E8M0 exponent tile -> cube ZZ scale fractal. Source is
436+// the RowMajor M̂xN exponent tile from TQuant<0,...>; equivalent to transposing it to
437+// NxM̂ then applying the stock ND->ZZ, but done in a single gather. M̂ must be even
438+// (M mod 64 == 0); see tquant-mxfp8-dn.md §5.5.
439+template <typename DstTileData, typename SrcTileData, typename TmpTileData>
440+__tf__ PTO_INTERNAL void TMovDnTo2Zz(typename DstTileData::TileDType __out__ dst,
441+ typename SrcTileData::TileDType __in__ src,
442+ typename TmpTileData::TileDType __in__ tmp, uint32_t hatM, uint32_t colsN)
443+{
444+ CommonCheckZZ<DstTileData, SrcTileData, TmpTileData>();
445+ static_assert(SrcTileData::isRowMajor && (SrcTileData::SFractal == SLayout::NoneBox),
446+ "TMov DN->ZZ: Source tile must be RowMajor with NoneBox layout.");
447+ static_assert(DstTileData::isRowMajor && (DstTileData::SFractal == SLayout::RowMajor),
448+ "TMov DN->ZZ: Destination Mat tile must use ColMajor + RowMajor fractal layout.");
449+ 
450+ __ubuf__ uint8_t *srcPtr = (__ubuf__ uint8_t *)__cce_get_tile_ptr(src);
451+ __ubuf__ uint8_t *dstPtr = (__ubuf__ uint8_t *)__cce_get_tile_ptr(dst);
452+ __ubuf__ uint8_t *tmpPtr = (__ubuf__ uint8_t *)__cce_get_tile_ptr(tmp);
453+ 
454+ __VEC_SCOPE__
455+ {
456+ GenerateB8IndicesDN2ZZToUB<DstTileData, SrcTileData>(dstPtr, srcPtr, tmpPtr, hatM, colsN);
457+ }
458+}
459+ 
380template <typename WorkT, typename SrcTileData>460template <typename WorkT, typename SrcTileData>
381PTO_INTERNAL void TMovNd2NzLoop(__ubuf__ WorkT *srcPtr, __ubuf__ WorkT *dstPtr, uint16_t repeatTimes,461PTO_INTERNAL void TMovNd2NzLoop(__ubuf__ WorkT *srcPtr, __ubuf__ WorkT *dstPtr, uint16_t repeatTimes,
382 uint16_t innerLoopNum, uint32_t validCol, uint32_t cfgVsstb, uint32_t cfgVsstbLast,462 uint16_t innerLoopNum, uint32_t validCol, uint32_t cfgVsstb, uint32_t cfgVsstbLast,
@@ -409,8 +489,11 @@ __tf__ PTO_INTERNAL void TMovToVecNd2Nz(typename DstTileData::TileDType __out__
409 static_assert((std::is_same<T, half>::value) || (std::is_same<T, bfloat16_t>::value) ||489 static_assert((std::is_same<T, half>::value) || (std::is_same<T, bfloat16_t>::value) ||
410 (std::is_same<T, float>::value) || (std::is_same<T, int32_t>::value) ||490 (std::is_same<T, float>::value) || (std::is_same<T, int32_t>::value) ||
411 (std::is_same<T, float8_e4m3_t>::value) || (std::is_same<T, float8_e5m2_t>::value) ||491 (std::is_same<T, float8_e4m3_t>::value) || (std::is_same<T, float8_e5m2_t>::value) ||
412- (std::is_same<T, hifloat8_t>::value) || (std::is_same<T, int8_t>::value),492+ (std::is_same<T, hifloat8_t>::value) || (std::is_same<T, int8_t>::value) ||
413- "Dst and src must be float/int32_t/half/bfloat16_t/int8_t/float8_e4m3_t/float8_e5m2_t/hifloat8_t.");493+ (std::is_same<T, uint8_t>::value) || (std::is_same<T, float4_e2m1x2_t>::value) ||
494+ (std::is_same<T, float4_e1m2x2_t>::value),
495+ "Dst and src must be float/int32_t/half/bfloat16_t/int8_t/uint8_t/float8_e4m3_t/float8_e5m2_t/"
496+ "hifloat8_t/float4_e2m1x2_t/float4_e1m2x2_t.");
414 __ubuf__ T *dstPtr = (__ubuf__ T *)__cce_get_tile_ptr(dst);497 __ubuf__ T *dstPtr = (__ubuf__ T *)__cce_get_tile_ptr(dst);
415 __ubuf__ T *srcPtr = (__ubuf__ T *)__cce_get_tile_ptr(src);498 __ubuf__ T *srcPtr = (__ubuf__ T *)__cce_get_tile_ptr(src);
416 constexpr int32_t srcRow = SrcTileData::Rows;499 constexpr int32_t srcRow = SrcTileData::Rows;
@@ -610,13 +693,22 @@ PTO_INTERNAL void TMOV_IMPL(DstTileData &dst, SrcTileData &src)
610 }693 }
611}694}
612 695 
613-template <typename DstTileData, typename SrcTileData, typename TmpTileData,696+// grp_axis-aware X->ZZ dispatch for the 3-arg TMOV(dst, src, tmp) form.
697+// grp_axis = 1 (default, ND): source is M x (N/32) exponents -> stock ND->ZZ.
698+// grp_axis = 0 (DN): source is (M/32) x N exponents (axis-0 grouping) -> DN->ZZ.
699+// Only the ZZ path is parameterised; all other TMOV overloads are unaffected.
700+template <int grp_axis = 1, typename DstTileData, typename SrcTileData, typename TmpTileData,
614 std::enable_if_t<(TmpTileData::Loc != TileType::Scaling), int> = 0>701 std::enable_if_t<(TmpTileData::Loc != TileType::Scaling), int> = 0>
615PTO_INTERNAL void TMOV_IMPL(DstTileData &dst, SrcTileData &src, TmpTileData &tmp)702PTO_INTERNAL void TMOV_IMPL(DstTileData &dst, SrcTileData &src, TmpTileData &tmp)
616{703{
617 CommonCheckZZ<DstTileData, SrcTileData, TmpTileData>();704 CommonCheckZZ<DstTileData, SrcTileData, TmpTileData>();
618- TMovNdTo2Zz<DstTileData, SrcTileData, TmpTileData>(dst.data(), src.data(), tmp.data(), dst.GetValidRow(),705+ if constexpr (grp_axis == 0) {
619- dst.GetValidCol());706+ TMovDnTo2Zz<DstTileData, SrcTileData, TmpTileData>(dst.data(), src.data(), tmp.data(), src.GetValidRow(),
707+ src.GetValidCol());
708+ } else {
709+ TMovNdTo2Zz<DstTileData, SrcTileData, TmpTileData>(dst.data(), src.data(), tmp.data(), dst.GetValidRow(),
710+ dst.GetValidCol());
711+ }
620}712}
621 713 
622template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, STPhase Phase = STPhase::Unspecified>714template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, STPhase Phase = STPhase::Unspecified>
@@ -78,7 +78,6 @@ set(CMAKE_CCE_COMPILE_OPTIONS
78 "SHELL:-mllvm -cce-aicore-addr-transform"78 "SHELL:-mllvm -cce-aicore-addr-transform"
79 "SHELL:-mllvm -cce-aicore-dcci-insert-for-scalar=false"79 "SHELL:-mllvm -cce-aicore-dcci-insert-for-scalar=false"
80)80)
81-set(CMAKE_CCE_COMPILE_OPTIONS "${CMAKE_CCE_COMPILE_OPTIONS} --cce-pto-enable")
82 81 
83if(DEBUG_MODE)82if(DEBUG_MODE)
84 message(STATUS "Debug Mode Enabled, Add Debug Options")83 message(STATUS "Debug Mode Enabled, Add Debug Options")
@@ -12,6 +12,8 @@
12 12 
13import math13import math
14import os14import os
15+from dataclasses import dataclass
16+from typing import Optional
15 17 
16import numpy as np18import numpy as np
17from ml_dtypes import bfloat16, float4_e2m1fn19from ml_dtypes import bfloat16, float4_e2m1fn
@@ -73,21 +75,14 @@ def fp32_to_e4m3(x):
73 75 
74 76 
75def nd2nz_mxfp8(data_fp8, m, n):77def nd2nz_mxfp8(data_fp8, m, n):
76- padded_rows16 = ((m + 15) // 16) * 1678+ # Stock ND->NZ for 1-byte data: [M,N] -> [n_groups, padded_m, 32]. No virtual_row+1
77- virtual_row = padded_rows16 + 179+ # (the +1 is a UB-internal stride; the GM NZ layout is plain [n_groups, padded_m, 32]).
78- padded_cols = ((n + 31) // 32) * 3280+ padded_m = ((m + 15) // 16) * 16
79- n_col_groups = padded_cols // 3281+ n_groups = ((n + 31) // 32) * 32 // 32
80- nz = np.zeros(virtual_row * padded_cols, dtype=np.int8)82+ reshaped = data_fp8.reshape(m, n_groups, 32) if data_fp8.ndim == 1 else data_fp8.reshape(m, n_groups, 32)
81- data_flat = data_fp8.reshape(-1) if data_fp8.ndim > 1 else data_fp883+ padded = np.zeros((padded_m, n_groups, 32), dtype=data_fp8.dtype)
82- for cg in range(n_col_groups):84+ padded[:m, :, :] = reshaped
83- for r in range(padded_rows16):85+ return np.transpose(padded, [1, 0, 2]).reshape(-1)
84- src_idx = r * padded_cols + cg * 32
85- dst_idx = cg * virtual_row * 32 + r * 32
86- if r < m:
87- nz[dst_idx : dst_idx + 32] = data_flat[src_idx : src_idx + 32]
88- else:
89- nz[dst_idx : dst_idx + 32] = 0
90- return nz
91 86 
92 87 
93def pack_e8_dn(e8m0, hat_m, n, padded_cols):88def pack_e8_dn(e8m0, hat_m, n, padded_cols):
@@ -100,8 +95,11 @@ def pack_e8_dn(e8m0, hat_m, n, padded_cols):
100 95 
101 96 
102def dn2zz_e8m0(e8m0_dn, hat_m, n):97def dn2zz_e8m0(e8m0_dn, hat_m, n):
103- # Row-major (ND) input -> ZZ is currently a flattened identity in this ST.98+ et = e8m0_dn.reshape(hat_m, n).T.copy()
atomgit-bot
atomgit-botatomgit-bot6月26日

🟡 Medium Priority

pack_e8_dn 返回 hat_m * padded_cols 元素的数组(包含尾部填充列),但 dn2zz_e8m0(e8_dn, hat_m, n) 按 (hat_m, n) reshape。当 n % 32 != 0 时 padded_cols > n,hat_m * padded_cols ≠ hat_m * n,reshape 会抛出 ValueError。旧代码 e8m0_dn[: hat_m * n].copy() 是安全的。当前所有测试用例的 n 均为 32 的倍数,因此未被触发,但对于非对齐输入会静默崩溃。

建议:在 dn2zz_e8m0 中,应先按 n 截取有效数据再进行 reshape:e8m0_dn = e8m0_dn[:hat_m * n],或者将 padded_cols 传入并在 reshape 中使用 padded_cols 然后切掉尾部填充列。最安全的方式是在函数开头加入 e8m0_dn = e8m0_dn[:hat_m * n].copy() 确保数组大小与 (hat_m, n) 匹配。

likedislike
不准确?
104- return e8m0_dn[: hat_m * n].copy()99+ rb = n // 16
100+ p = hat_m // 2
101+ zz = et.reshape(rb, 16, p, 2).transpose(0, 2, 1, 3).reshape(-1).astype(np.uint8)
102+ return zz
105 103 
106 104 
107def quant_bf16_to_mxfp8_dn(src_bf16_fp32, m, n_pad):105def quant_bf16_to_mxfp8_dn(src_bf16_fp32, m, n_pad):
@@ -310,16 +308,41 @@ def _gen_src(m, n_pad):
310 return base * group_max_repeated * 10000.0308 return base * group_max_repeated * 10000.0
311 309 
312 310 
313-def _write_golden(out_dir, input_bytes, fp8_nd, e8_dn, group_max_bytes):311+@dataclass
312+class GoldenDataFP8:
313+ input_bytes: bytes
314+ fp8_nd: np.ndarray
315+ e8_dn: np.ndarray
316+ group_max_bytes: bytes
317+ fp8_nz: Optional[np.ndarray] = None
318+ e8_zz: Optional[np.ndarray] = None
319+ 
320+ 
321+@dataclass
322+class GoldenDataFP4:
323+ input_bytes: bytes
324+ fp4_nd: np.ndarray
325+ e8_dn: np.ndarray
326+ fp4_nz: np.ndarray
327+ group_max_bytes: bytes
328+ 
329+ 
330+def _write_golden(out_dir, golden):
314 os.makedirs(out_dir, exist_ok=True)331 os.makedirs(out_dir, exist_ok=True)
315 with open(os.path.join(out_dir, "input.bin"), "wb") as f:332 with open(os.path.join(out_dir, "input.bin"), "wb") as f:
316- f.write(input_bytes)333+ f.write(golden.input_bytes)
317 with open(os.path.join(out_dir, "golden_fp8_nd.bin"), "wb") as f:334 with open(os.path.join(out_dir, "golden_fp8_nd.bin"), "wb") as f:
318- f.write(fp8_nd.tobytes())335+ f.write(golden.fp8_nd.tobytes())
319 with open(os.path.join(out_dir, "golden_e8_dn.bin"), "wb") as f:336 with open(os.path.join(out_dir, "golden_e8_dn.bin"), "wb") as f:
320- f.write(e8_dn.tobytes())337+ f.write(golden.e8_dn.tobytes())
321 with open(os.path.join(out_dir, "golden_group_max.bin"), "wb") as f:338 with open(os.path.join(out_dir, "golden_group_max.bin"), "wb") as f:
322- f.write(group_max_bytes)339+ f.write(golden.group_max_bytes)
340+ if golden.fp8_nz is not None:
341+ with open(os.path.join(out_dir, "golden_fp8_nz.bin"), "wb") as f:
342+ f.write(golden.fp8_nz.tobytes())
343+ if golden.e8_zz is not None:
344+ with open(os.path.join(out_dir, "golden_e8_zz.bin"), "wb") as f:
345+ f.write(golden.e8_zz.tobytes())
323 346 
324 347 
325def gen_golden_data(case_name, m, n):348def gen_golden_data(case_name, m, n):
@@ -328,39 +351,57 @@ def gen_golden_data(case_name, m, n):
328 bf16_bits = fp32_to_bf16_bits(src).reshape(m, n_pad)351 bf16_bits = fp32_to_bf16_bits(src).reshape(m, n_pad)
329 src_bf16_fp32 = bf16_bits_to_fp32(bf16_bits.flatten()).reshape(m, n_pad)352 src_bf16_fp32 = bf16_bits_to_fp32(bf16_bits.flatten()).reshape(m, n_pad)
330 353 
331- fp8_nd, e8_dn, _, _ = quant_bf16_to_mxfp8_dn(src_bf16_fp32, m, n_pad)354+ fp8_nd, e8_dn, fp8_nz, e8_zz = quant_bf16_to_mxfp8_dn(src_bf16_fp32, m, n_pad)
332 355 
333 group_max = get_group_max_dn(src_bf16_fp32, group_size=32)356 group_max = get_group_max_dn(src_bf16_fp32, group_size=32)
334 golden_group_max_bf16 = fp32_to_bf16_bits(group_max)357 golden_group_max_bf16 = fp32_to_bf16_bits(group_max)
335 358 
336 out_dir = os.path.join(GOLDEN_DIR, case_name)359 out_dir = os.path.join(GOLDEN_DIR, case_name)
337- _write_golden(out_dir, bf16_bits.reshape(-1).tobytes(), fp8_nd, e8_dn, golden_group_max_bf16.reshape(-1).tobytes())360+ golden = GoldenDataFP8(
361+ input_bytes=bf16_bits.reshape(-1).tobytes(),
362+ fp8_nd=fp8_nd,
363+ e8_dn=e8_dn,
364+ group_max_bytes=golden_group_max_bf16.reshape(-1).tobytes(),
365+ fp8_nz=fp8_nz,
366+ e8_zz=e8_zz,
367+ )
368+ _write_golden(out_dir, golden)
338 369 
339 370 
340def gen_golden_data_fp32(case_name, m, n):371def gen_golden_data_fp32(case_name, m, n):
341 n_pad = n372 n_pad = n
342 src = _gen_src(m, n_pad)373 src = _gen_src(m, n_pad)
343 374 
344- fp8_nd, e8_dn, _, _ = quant_bf16_to_mxfp8_dn(src, m, n_pad)375+ fp8_nd, e8_dn, fp8_nz, e8_zz = quant_bf16_to_mxfp8_dn(src, m, n_pad)
345 376 
346 group_max = get_group_max_dn(src, group_size=32)377 group_max = get_group_max_dn(src, group_size=32)
347 golden_group_max_f32 = group_max.astype(np.float32).view(np.uint32)378 golden_group_max_f32 = group_max.astype(np.float32).view(np.uint32)
348 379 
349 out_dir = os.path.join(GOLDEN_DIR, case_name)380 out_dir = os.path.join(GOLDEN_DIR, case_name)
350 input_bytes = src.astype(np.float32).view(np.uint32).reshape(-1).tobytes()381 input_bytes = src.astype(np.float32).view(np.uint32).reshape(-1).tobytes()
351- _write_golden(out_dir, input_bytes, fp8_nd, e8_dn, golden_group_max_f32.reshape(-1).tobytes())382+ golden = GoldenDataFP8(
383+ input_bytes=input_bytes,
384+ fp8_nd=fp8_nd,
385+ e8_dn=e8_dn,
386+ group_max_bytes=golden_group_max_f32.reshape(-1).tobytes(),
387+ fp8_nz=fp8_nz,
388+ e8_zz=e8_zz,
389+ )
390+ _write_golden(out_dir, golden)
352 391 
353 392 
354-def _write_golden_mxfp4(out_dir, input_bytes, fp4_nd, e8_dn, group_max_bytes):393+def _write_golden_mxfp4(out_dir, golden):
355 os.makedirs(out_dir, exist_ok=True)394 os.makedirs(out_dir, exist_ok=True)
356 with open(os.path.join(out_dir, "input.bin"), "wb") as f:395 with open(os.path.join(out_dir, "input.bin"), "wb") as f:
357- f.write(input_bytes)396+ f.write(golden.input_bytes)
358 with open(os.path.join(out_dir, "golden_fp4_nd.bin"), "wb") as f:397 with open(os.path.join(out_dir, "golden_fp4_nd.bin"), "wb") as f:
359- f.write(fp4_nd.tobytes())398+ f.write(golden.fp4_nd.tobytes())
360 with open(os.path.join(out_dir, "golden_e8_dn.bin"), "wb") as f:399 with open(os.path.join(out_dir, "golden_e8_dn.bin"), "wb") as f:
361- f.write(e8_dn.tobytes())400+ f.write(golden.e8_dn.tobytes())
401+ with open(os.path.join(out_dir, "golden_fp4_nz.bin"), "wb") as f:
402+ f.write(golden.fp4_nz.tobytes())
362 with open(os.path.join(out_dir, "golden_group_max.bin"), "wb") as f:403 with open(os.path.join(out_dir, "golden_group_max.bin"), "wb") as f:
363- f.write(group_max_bytes)404+ f.write(golden.group_max_bytes)
364 405 
365 406 
366def gen_golden_data_mxfp4_bf16(case_name, m, n):407def gen_golden_data_mxfp4_bf16(case_name, m, n):
@@ -370,12 +411,20 @@ def gen_golden_data_mxfp4_bf16(case_name, m, n):
370 src_bf16_fp32 = bf16_bits_to_fp32(bf16_bits.flatten()).reshape(m, n_pad)411 src_bf16_fp32 = bf16_bits_to_fp32(bf16_bits.flatten()).reshape(m, n_pad)
371 412 
372 fp4_nd, e8_dn, group_max = quant_bf16_to_mxfp4_dn(src_bf16_fp32, m, n_pad)413 fp4_nd, e8_dn, group_max = quant_bf16_to_mxfp4_dn(src_bf16_fp32, m, n_pad)
414+ fp4_padded = np.zeros((m, n_pad // 2), dtype=np.uint8)
415+ fp4_padded[:, : n_pad // 2] = fp4_nd.reshape(m, n_pad // 2)
416+ fp4_nz = nd2nz_mxfp8(fp4_padded, m, n_pad // 2)
373 golden_group_max_bf16 = fp32_to_bf16_bits(group_max)417 golden_group_max_bf16 = fp32_to_bf16_bits(group_max)
374 418 
375 out_dir = os.path.join(GOLDEN_DIR, case_name)419 out_dir = os.path.join(GOLDEN_DIR, case_name)
376- _write_golden_mxfp4(420+ golden = GoldenDataFP4(
377- out_dir, bf16_bits.reshape(-1).tobytes(), fp4_nd, e8_dn, golden_group_max_bf16.reshape(-1).tobytes()421+ input_bytes=bf16_bits.reshape(-1).tobytes(),
422+ fp4_nd=fp4_nd,
423+ e8_dn=e8_dn,
424+ fp4_nz=fp4_nz,
425+ group_max_bytes=golden_group_max_bf16.reshape(-1).tobytes(),
378 )426 )
427+ _write_golden_mxfp4(out_dir, golden)
379 428 
380 429 
381def _gen_src_fp16_safe(m, n_pad):430def _gen_src_fp16_safe(m, n_pad):
@@ -400,14 +449,22 @@ def gen_golden_data_mxfp4_fp16(case_name, m, n):
400 src_fp16_bits = src.view(np.uint16)449 src_fp16_bits = src.view(np.uint16)
401 450 
402 fp4_nd, e8_dn, group_max = quant_fp16_to_mxfp4_dn(src, m, n_pad)451 fp4_nd, e8_dn, group_max = quant_fp16_to_mxfp4_dn(src, m, n_pad)
452+ fp4_padded = np.zeros((m, n_pad // 2), dtype=np.uint8)
453+ fp4_padded[:, : n_pad // 2] = fp4_nd.reshape(m, n_pad // 2)
454+ fp4_nz = nd2nz_mxfp8(fp4_padded, m, n_pad // 2)
403 # DN reduces fp16 abs directly and stores the max as fp16 bits (kernel MaxTile455 # DN reduces fp16 abs directly and stores the max as fp16 bits (kernel MaxTile
404 # is half when T=half), so the golden group_max is fp16 bits too.456 # is half when T=half), so the golden group_max is fp16 bits too.
405 golden_group_max_fp16 = group_max.view(np.uint16)457 golden_group_max_fp16 = group_max.view(np.uint16)
406 458 
407 out_dir = os.path.join(GOLDEN_DIR, case_name)459 out_dir = os.path.join(GOLDEN_DIR, case_name)
408- _write_golden_mxfp4(460+ golden = GoldenDataFP4(
409- out_dir, src_fp16_bits.reshape(-1).tobytes(), fp4_nd, e8_dn, golden_group_max_fp16.reshape(-1).tobytes()461+ input_bytes=src_fp16_bits.reshape(-1).tobytes(),
462+ fp4_nd=fp4_nd,
463+ e8_dn=e8_dn,
464+ fp4_nz=fp4_nz,
465+ group_max_bytes=golden_group_max_fp16.reshape(-1).tobytes(),
410 )466 )
467+ _write_golden_mxfp4(out_dir, golden)
411 468 
412 469 
413if __name__ == "__main__":470if __name__ == "__main__":
@@ -17,19 +17,21 @@ using namespace PtoTestCommon;
17 17 
18namespace TQuantDNTest {18namespace TQuantDNTest {
19 19 
20-template <int Stage, int M, int N, int N_pad>20+template <int M, int N, int N_pad>
21void LaunchTQuantDN(uint16_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, uint16_t *max_dn,21void LaunchTQuantDN(uint16_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, uint16_t *max_dn,
22 void *stream);22 void *stream);
23 23 
24-template <int Stage, int M, int N, int N_pad>24+template <int M, int N, int N_pad>
25void LaunchTQuantDN_fp32(uint32_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz,25void LaunchTQuantDN_fp32(uint32_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz,
26 uint32_t *max_dn, void *stream);26 uint32_t *max_dn, void *stream);
27 27 
28template <int M, int N, int N_pad>28template <int M, int N, int N_pad>
29-void LaunchTQuantDN_MXFP4_bf16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint16_t *max_dn, void *stream);29+void LaunchTQuantDN_MXFP4_bf16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint8_t *fp4_nz, uint16_t *max_dn,
30+ void *stream);
30 31 
31template <int M, int N, int N_pad>32template <int M, int N, int N_pad>
32-void LaunchTQuantDN_MXFP4_fp16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint16_t *max_dn, void *stream);33+void LaunchTQuantDN_MXFP4_fp16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint8_t *fp4_nz, uint16_t *max_dn,
34+ void *stream);
33 35 
34} // namespace TQuantDNTest36} // namespace TQuantDNTest
35 37 
@@ -72,7 +74,7 @@ void test_tquant_dn_bf16()
72 size_t srcFileSize = M * paddedCols * sizeof(uint16_t);74 size_t srcFileSize = M * paddedCols * sizeof(uint16_t);
73 size_t fp8NDFileSize = M * paddedCols * sizeof(int8_t);75 size_t fp8NDFileSize = M * paddedCols * sizeof(int8_t);
74 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);76 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);
75- size_t fp8NZFileSize = virtualRow * paddedCols * sizeof(int8_t);77+ size_t fp8NZFileSize = M * paddedCols * sizeof(int8_t);
76 size_t e8ZZFileSize = numGroupsFlatAligned * sizeof(uint8_t);78 size_t e8ZZFileSize = numGroupsFlatAligned * sizeof(uint8_t);
77 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint16_t);79 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint16_t);
78 80 
@@ -113,34 +115,48 @@ void test_tquant_dn_bf16()
113 115 
114 const std::string goldenDir = GetGoldenDir();116 const std::string goldenDir = GetGoldenDir();
115 117 
116- // Stage 1: after TQUANT — FP8 ND + E8M0 DN + per-group max118+ // Full DN pipeline: TQuant(DN) + TMOV(ND->NZ) + TMOV<0>(DN->ZZ).
117- TQuantDNTest::LaunchTQuantDN<1, M, N, N_pad>(srcDevice, fp8NDDevice, e8DNDevice, nullptr, nullptr, maxDNDevice,119+ TQuantDNTest::LaunchTQuantDN<M, N, N_pad>(srcDevice, fp8NDDevice, e8DNDevice, fp8NZDevice, e8ZZDevice, maxDNDevice,
118- stream);120+ stream);
119 aclError syncRet = aclrtSynchronizeStream(stream);121 aclError syncRet = aclrtSynchronizeStream(stream);
120- ASSERT_EQ(syncRet, ACL_SUCCESS) << "Stage1 sync failed: " << aclGetRecentErrMsg();122+ ASSERT_EQ(syncRet, ACL_SUCCESS) << "DN pipeline sync failed: " << aclGetRecentErrMsg();
121 123 
122 aclrtMemcpy(fp8NDHost, fp8NDFileSize, fp8NDDevice, fp8NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);124 aclrtMemcpy(fp8NDHost, fp8NDFileSize, fp8NDDevice, fp8NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
123 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);125 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
126+ aclrtMemcpy(fp8NZHost, fp8NZFileSize, fp8NZDevice, fp8NZFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
127+ aclrtMemcpy(e8ZZHost, e8ZZFileSize, e8ZZDevice, e8ZZFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
124 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);128 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
125 WriteFile(goldenDir + "/output_fp8_nd.bin", fp8NDHost, fp8NDFileSize);129 WriteFile(goldenDir + "/output_fp8_nd.bin", fp8NDHost, fp8NDFileSize);
126 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);130 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);
131+ WriteFile(goldenDir + "/output_fp8_nz.bin", fp8NZHost, fp8NZFileSize);
132+ WriteFile(goldenDir + "/output_e8_zz.bin", e8ZZHost, e8ZZFileSize);
127 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);133 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);
128 134 
129 std::vector<uint8_t> goldenFp8Nd(fp8NDFileSize);135 std::vector<uint8_t> goldenFp8Nd(fp8NDFileSize);
130 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);136 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);
137+ std::vector<uint8_t> goldenFp8Nz(fp8NZFileSize);
138+ std::vector<uint8_t> goldenE8Zz(e8ZZFileSize);
131 std::vector<uint16_t> goldenGroupMax(maxDNFileSize / sizeof(uint16_t));139 std::vector<uint16_t> goldenGroupMax(maxDNFileSize / sizeof(uint16_t));
132 std::vector<uint8_t> outFp8Nd(fp8NDFileSize);140 std::vector<uint8_t> outFp8Nd(fp8NDFileSize);
133 std::vector<uint8_t> outE8Dn(e8DNFileSize);141 std::vector<uint8_t> outE8Dn(e8DNFileSize);
142+ std::vector<uint8_t> outFp8Nz(fp8NZFileSize);
143+ std::vector<uint8_t> outE8Zz(e8ZZFileSize);
134 std::vector<uint16_t> outGroupMax(maxDNFileSize / sizeof(uint16_t));144 std::vector<uint16_t> outGroupMax(maxDNFileSize / sizeof(uint16_t));
135 ReadFile(goldenDir + "/golden_fp8_nd.bin", fp8NDFileSize, goldenFp8Nd.data(), fp8NDFileSize);145 ReadFile(goldenDir + "/golden_fp8_nd.bin", fp8NDFileSize, goldenFp8Nd.data(), fp8NDFileSize);
136 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);146 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);
147+ ReadFile(goldenDir + "/golden_fp8_nz.bin", fp8NZFileSize, goldenFp8Nz.data(), fp8NZFileSize);
148+ ReadFile(goldenDir + "/golden_e8_zz.bin", e8ZZFileSize, goldenE8Zz.data(), e8ZZFileSize);
137 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);149 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);
138 ReadFile(goldenDir + "/output_fp8_nd.bin", fp8NDFileSize, outFp8Nd.data(), fp8NDFileSize);150 ReadFile(goldenDir + "/output_fp8_nd.bin", fp8NDFileSize, outFp8Nd.data(), fp8NDFileSize);
139 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);151 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);
152+ ReadFile(goldenDir + "/output_fp8_nz.bin", fp8NZFileSize, outFp8Nz.data(), fp8NZFileSize);
153+ ReadFile(goldenDir + "/output_e8_zz.bin", e8ZZFileSize, outE8Zz.data(), e8ZZFileSize);
140 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);154 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);
141- ExpectGoldenMatch("Stage1_AfterTQuant", "fp8_nd", goldenFp8Nd, outFp8Nd);155+ ExpectGoldenMatch("DN_Pipeline", "fp8_nd", goldenFp8Nd, outFp8Nd);
142- ExpectGoldenMatch("Stage1_AfterTQuant", "e8_dn (exponents)", goldenE8Dn, outE8Dn);156+ ExpectGoldenMatch("DN_Pipeline", "e8_dn (exponents)", goldenE8Dn, outE8Dn);
143- ExpectGoldenMatch("Stage1_AfterTQuant", "group_max", goldenGroupMax, outGroupMax);157+ ExpectGoldenMatch("DN_Pipeline", "fp8_nz", goldenFp8Nz, outFp8Nz);
158+ ExpectGoldenMatch("DN_Pipeline", "e8_zz", goldenE8Zz, outE8Zz);
159+ ExpectGoldenMatch("DN_Pipeline", "group_max", goldenGroupMax, outGroupMax);
144 160 
145 aclrtFree(srcDevice);161 aclrtFree(srcDevice);
146 aclrtFree(fp8NDDevice);162 aclrtFree(fp8NDDevice);
@@ -206,7 +222,7 @@ void test_tquant_dn_fp32()
206 size_t srcFileSize = M * paddedCols * sizeof(uint32_t);222 size_t srcFileSize = M * paddedCols * sizeof(uint32_t);
207 size_t fp8NDFileSize = M * paddedCols * sizeof(int8_t);223 size_t fp8NDFileSize = M * paddedCols * sizeof(int8_t);
208 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);224 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);
209- size_t fp8NZFileSize = virtualRow * paddedCols * sizeof(int8_t);225+ size_t fp8NZFileSize = M * paddedCols * sizeof(int8_t);
210 size_t e8ZZFileSize = numGroupsFlatAligned * sizeof(uint8_t);226 size_t e8ZZFileSize = numGroupsFlatAligned * sizeof(uint8_t);
211 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint32_t);227 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint32_t);
212 228 
@@ -247,33 +263,48 @@ void test_tquant_dn_fp32()
247 263 
248 const std::string goldenDir = GetGoldenDir();264 const std::string goldenDir = GetGoldenDir();
249 265 
250- TQuantDNTest::LaunchTQuantDN_fp32<1, M, N, N_pad>(srcDevice, fp8NDDevice, e8DNDevice, nullptr, nullptr, maxDNDevice,266+ // Full DN pipeline: TQuant(DN) + TMOV(ND->NZ) + TMOV<0>(DN->ZZ).
251- stream);267+ TQuantDNTest::LaunchTQuantDN_fp32<M, N, N_pad>(srcDevice, fp8NDDevice, e8DNDevice, fp8NZDevice, e8ZZDevice,
268+ maxDNDevice, stream);
252 aclError syncRet = aclrtSynchronizeStream(stream);269 aclError syncRet = aclrtSynchronizeStream(stream);
253- ASSERT_EQ(syncRet, ACL_SUCCESS) << "Stage1 sync failed: " << aclGetRecentErrMsg();270+ ASSERT_EQ(syncRet, ACL_SUCCESS) << "DN pipeline fp32 sync failed: " << aclGetRecentErrMsg();
254 271 
255 aclrtMemcpy(fp8NDHost, fp8NDFileSize, fp8NDDevice, fp8NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);272 aclrtMemcpy(fp8NDHost, fp8NDFileSize, fp8NDDevice, fp8NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
256 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);273 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
274+ aclrtMemcpy(fp8NZHost, fp8NZFileSize, fp8NZDevice, fp8NZFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
275+ aclrtMemcpy(e8ZZHost, e8ZZFileSize, e8ZZDevice, e8ZZFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
257 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);276 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
258 WriteFile(goldenDir + "/output_fp8_nd.bin", fp8NDHost, fp8NDFileSize);277 WriteFile(goldenDir + "/output_fp8_nd.bin", fp8NDHost, fp8NDFileSize);
259 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);278 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);
279+ WriteFile(goldenDir + "/output_fp8_nz.bin", fp8NZHost, fp8NZFileSize);
280+ WriteFile(goldenDir + "/output_e8_zz.bin", e8ZZHost, e8ZZFileSize);
260 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);281 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);
261 282 
262 std::vector<uint8_t> goldenFp8Nd(fp8NDFileSize);283 std::vector<uint8_t> goldenFp8Nd(fp8NDFileSize);
263 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);284 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);
285+ std::vector<uint8_t> goldenFp8Nz(fp8NZFileSize);
286+ std::vector<uint8_t> goldenE8Zz(e8ZZFileSize);
264 std::vector<uint32_t> goldenGroupMax(maxDNFileSize / sizeof(uint32_t));287 std::vector<uint32_t> goldenGroupMax(maxDNFileSize / sizeof(uint32_t));
265 std::vector<uint8_t> outFp8Nd(fp8NDFileSize);288 std::vector<uint8_t> outFp8Nd(fp8NDFileSize);
266 std::vector<uint8_t> outE8Dn(e8DNFileSize);289 std::vector<uint8_t> outE8Dn(e8DNFileSize);
290+ std::vector<uint8_t> outFp8Nz(fp8NZFileSize);
291+ std::vector<uint8_t> outE8Zz(e8ZZFileSize);
267 std::vector<uint32_t> outGroupMax(maxDNFileSize / sizeof(uint32_t));292 std::vector<uint32_t> outGroupMax(maxDNFileSize / sizeof(uint32_t));
268 ReadFile(goldenDir + "/golden_fp8_nd.bin", fp8NDFileSize, goldenFp8Nd.data(), fp8NDFileSize);293 ReadFile(goldenDir + "/golden_fp8_nd.bin", fp8NDFileSize, goldenFp8Nd.data(), fp8NDFileSize);
269 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);294 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);
295+ ReadFile(goldenDir + "/golden_fp8_nz.bin", fp8NZFileSize, goldenFp8Nz.data(), fp8NZFileSize);
296+ ReadFile(goldenDir + "/golden_e8_zz.bin", e8ZZFileSize, goldenE8Zz.data(), e8ZZFileSize);
270 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);297 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);
271 ReadFile(goldenDir + "/output_fp8_nd.bin", fp8NDFileSize, outFp8Nd.data(), fp8NDFileSize);298 ReadFile(goldenDir + "/output_fp8_nd.bin", fp8NDFileSize, outFp8Nd.data(), fp8NDFileSize);
272 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);299 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);
300+ ReadFile(goldenDir + "/output_fp8_nz.bin", fp8NZFileSize, outFp8Nz.data(), fp8NZFileSize);
301+ ReadFile(goldenDir + "/output_e8_zz.bin", e8ZZFileSize, outE8Zz.data(), e8ZZFileSize);
273 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);302 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);
274- ExpectGoldenMatch("Stage1_AfterTQuant", "fp8_nd", goldenFp8Nd, outFp8Nd);303+ ExpectGoldenMatch("DN_Pipeline_fp32", "fp8_nd", goldenFp8Nd, outFp8Nd);
275- ExpectGoldenMatch("Stage1_AfterTQuant", "e8_dn (exponents)", goldenE8Dn, outE8Dn);304+ ExpectGoldenMatch("DN_Pipeline_fp32", "e8_dn (exponents)", goldenE8Dn, outE8Dn);
276- ExpectGoldenMatch("Stage1_AfterTQuant", "group_max", goldenGroupMax, outGroupMax);305+ ExpectGoldenMatch("DN_Pipeline_fp32", "fp8_nz", goldenFp8Nz, outFp8Nz);
306+ ExpectGoldenMatch("DN_Pipeline_fp32", "e8_zz", goldenE8Zz, outE8Zz);
307+ ExpectGoldenMatch("DN_Pipeline_fp32", "group_max", goldenGroupMax, outGroupMax);
277 308 
278 aclrtFree(srcDevice);309 aclrtFree(srcDevice);
279 aclrtFree(fp8NDDevice);310 aclrtFree(fp8NDDevice);
@@ -307,6 +338,7 @@ TEST_F(TQUANTDNTest, case_fp32_64x256)
307 338 
308// MXFP4 (E2M1) DN tests. Both bf16 and fp16 sources share the same UB/GM shape:339// MXFP4 (E2M1) DN tests. Both bf16 and fp16 sources share the same UB/GM shape:
309// fp4_nd : M * packedCols bytes (packedCols = paddedCols/2)340// fp4_nd : M * packedCols bytes (packedCols = paddedCols/2)
341+// fp4_nz : M * packedCols bytes (ND->NZ packed FP4, block = 32B = 64 FP4 values)
310// e8_dn : hatM * paddedCols bytes342// e8_dn : hatM * paddedCols bytes
311// max_dn : hatM * paddedCols * sizeof(b16) bytes343// max_dn : hatM * paddedCols * sizeof(b16) bytes
312template <int M, int N, int N_pad>344template <int M, int N, int N_pad>
@@ -318,6 +350,7 @@ void test_tquant_dn_mxfp4_bf16()
318 constexpr int packedCols = paddedCols / 2;350 constexpr int packedCols = paddedCols / 2;
319 size_t srcFileSize = M * paddedCols * sizeof(uint16_t);351 size_t srcFileSize = M * paddedCols * sizeof(uint16_t);
320 size_t fp4NDFileSize = M * packedCols * sizeof(uint8_t);352 size_t fp4NDFileSize = M * packedCols * sizeof(uint8_t);
353+ size_t fp4NZFileSize = M * packedCols * sizeof(uint8_t);
321 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);354 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);
322 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint16_t);355 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint16_t);
323 356 
@@ -328,20 +361,24 @@ void test_tquant_dn_mxfp4_bf16()
328 361 
329 uint8_t *srcHost;362 uint8_t *srcHost;
330 uint8_t *fp4NDHost;363 uint8_t *fp4NDHost;
364+ uint8_t *fp4NZHost;
331 uint8_t *e8DNHost;365 uint8_t *e8DNHost;
332 uint16_t *maxDNHost;366 uint16_t *maxDNHost;
333 uint8_t *srcDevice;367 uint8_t *srcDevice;
334 uint8_t *fp4NDDevice;368 uint8_t *fp4NDDevice;
369+ uint8_t *fp4NZDevice;
335 uint8_t *e8DNDevice;370 uint8_t *e8DNDevice;
336 uint16_t *maxDNDevice;371 uint16_t *maxDNDevice;
337 372 
338 aclrtMallocHost((void **)(&srcHost), srcFileSize);373 aclrtMallocHost((void **)(&srcHost), srcFileSize);
339 aclrtMallocHost((void **)(&fp4NDHost), fp4NDFileSize);374 aclrtMallocHost((void **)(&fp4NDHost), fp4NDFileSize);
375+ aclrtMallocHost((void **)(&fp4NZHost), fp4NZFileSize);
340 aclrtMallocHost((void **)(&e8DNHost), e8DNFileSize);376 aclrtMallocHost((void **)(&e8DNHost), e8DNFileSize);
341 aclrtMallocHost((void **)(&maxDNHost), maxDNFileSize);377 aclrtMallocHost((void **)(&maxDNHost), maxDNFileSize);
342 378 
343 aclrtMalloc((void **)&srcDevice, srcFileSize, ACL_MEM_MALLOC_HUGE_FIRST);379 aclrtMalloc((void **)&srcDevice, srcFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
344 aclrtMalloc((void **)&fp4NDDevice, fp4NDFileSize, ACL_MEM_MALLOC_HUGE_FIRST);380 aclrtMalloc((void **)&fp4NDDevice, fp4NDFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
381+ aclrtMalloc((void **)&fp4NZDevice, fp4NZFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
345 aclrtMalloc((void **)&e8DNDevice, e8DNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);382 aclrtMalloc((void **)&e8DNDevice, e8DNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
346 aclrtMalloc((void **)&maxDNDevice, maxDNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);383 aclrtMalloc((void **)&maxDNDevice, maxDNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
347 384 
@@ -350,40 +387,49 @@ void test_tquant_dn_mxfp4_bf16()
350 387 
351 const std::string goldenDir = GetGoldenDir();388 const std::string goldenDir = GetGoldenDir();
352 389 
353- TQuantDNTest::LaunchTQuantDN_MXFP4_bf16<M, N, N_pad>((uint16_t *)srcDevice, fp4NDDevice, e8DNDevice, maxDNDevice,390+ TQuantDNTest::LaunchTQuantDN_MXFP4_bf16<M, N, N_pad>((uint16_t *)srcDevice, fp4NDDevice, e8DNDevice, fp4NZDevice,
354- stream);391+ maxDNDevice, stream);
355 aclError syncRet = aclrtSynchronizeStream(stream);392 aclError syncRet = aclrtSynchronizeStream(stream);
356 ASSERT_EQ(syncRet, ACL_SUCCESS) << "MXFP4 bf16 DN sync failed: " << aclGetRecentErrMsg();393 ASSERT_EQ(syncRet, ACL_SUCCESS) << "MXFP4 bf16 DN sync failed: " << aclGetRecentErrMsg();
357 394 
358 aclrtMemcpy(fp4NDHost, fp4NDFileSize, fp4NDDevice, fp4NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);395 aclrtMemcpy(fp4NDHost, fp4NDFileSize, fp4NDDevice, fp4NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
396+ aclrtMemcpy(fp4NZHost, fp4NZFileSize, fp4NZDevice, fp4NZFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
359 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);397 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
360 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);398 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
361 WriteFile(goldenDir + "/output_fp4_nd.bin", fp4NDHost, fp4NDFileSize);399 WriteFile(goldenDir + "/output_fp4_nd.bin", fp4NDHost, fp4NDFileSize);
400+ WriteFile(goldenDir + "/output_fp4_nz.bin", fp4NZHost, fp4NZFileSize);
362 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);401 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);
363 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);402 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);
364 403 
365 std::vector<uint8_t> goldenFp4Nd(fp4NDFileSize);404 std::vector<uint8_t> goldenFp4Nd(fp4NDFileSize);
405+ std::vector<uint8_t> goldenFp4Nz(fp4NZFileSize);
366 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);406 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);
367 std::vector<uint16_t> goldenGroupMax(maxDNFileSize / sizeof(uint16_t));407 std::vector<uint16_t> goldenGroupMax(maxDNFileSize / sizeof(uint16_t));
368 std::vector<uint8_t> outFp4Nd(fp4NDFileSize);408 std::vector<uint8_t> outFp4Nd(fp4NDFileSize);
409+ std::vector<uint8_t> outFp4Nz(fp4NZFileSize);
369 std::vector<uint8_t> outE8Dn(e8DNFileSize);410 std::vector<uint8_t> outE8Dn(e8DNFileSize);
370 std::vector<uint16_t> outGroupMax(maxDNFileSize / sizeof(uint16_t));411 std::vector<uint16_t> outGroupMax(maxDNFileSize / sizeof(uint16_t));
371 ReadFile(goldenDir + "/golden_fp4_nd.bin", fp4NDFileSize, goldenFp4Nd.data(), fp4NDFileSize);412 ReadFile(goldenDir + "/golden_fp4_nd.bin", fp4NDFileSize, goldenFp4Nd.data(), fp4NDFileSize);
413+ ReadFile(goldenDir + "/golden_fp4_nz.bin", fp4NZFileSize, goldenFp4Nz.data(), fp4NZFileSize);
372 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);414 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);
373 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);415 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);
374 ReadFile(goldenDir + "/output_fp4_nd.bin", fp4NDFileSize, outFp4Nd.data(), fp4NDFileSize);416 ReadFile(goldenDir + "/output_fp4_nd.bin", fp4NDFileSize, outFp4Nd.data(), fp4NDFileSize);
417+ ReadFile(goldenDir + "/output_fp4_nz.bin", fp4NZFileSize, outFp4Nz.data(), fp4NZFileSize);
375 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);418 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);
376 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);419 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);
377 ExpectGoldenMatch("MXFP4_BF16_DN", "fp4_nd", goldenFp4Nd, outFp4Nd);420 ExpectGoldenMatch("MXFP4_BF16_DN", "fp4_nd", goldenFp4Nd, outFp4Nd);
421+ ExpectGoldenMatch("MXFP4_BF16_DN", "fp4_nz", goldenFp4Nz, outFp4Nz);
378 ExpectGoldenMatch("MXFP4_BF16_DN", "e8_dn (exponents)", goldenE8Dn, outE8Dn);422 ExpectGoldenMatch("MXFP4_BF16_DN", "e8_dn (exponents)", goldenE8Dn, outE8Dn);
379 ExpectGoldenMatch("MXFP4_BF16_DN", "group_max", goldenGroupMax, outGroupMax);423 ExpectGoldenMatch("MXFP4_BF16_DN", "group_max", goldenGroupMax, outGroupMax);
380 424 
381 aclrtFree(srcDevice);425 aclrtFree(srcDevice);
382 aclrtFree(fp4NDDevice);426 aclrtFree(fp4NDDevice);
427+ aclrtFree(fp4NZDevice);
383 aclrtFree(e8DNDevice);428 aclrtFree(e8DNDevice);
384 aclrtFree(maxDNDevice);429 aclrtFree(maxDNDevice);
385 aclrtFreeHost(srcHost);430 aclrtFreeHost(srcHost);
386 aclrtFreeHost(fp4NDHost);431 aclrtFreeHost(fp4NDHost);
432+ aclrtFreeHost(fp4NZHost);
387 aclrtFreeHost(e8DNHost);433 aclrtFreeHost(e8DNHost);
388 aclrtFreeHost(maxDNHost);434 aclrtFreeHost(maxDNHost);
389 aclrtDestroyStream(stream);435 aclrtDestroyStream(stream);
@@ -413,6 +459,7 @@ void test_tquant_dn_mxfp4_fp16()
413 constexpr int packedCols = paddedCols / 2;459 constexpr int packedCols = paddedCols / 2;
414 size_t srcFileSize = M * paddedCols * sizeof(uint16_t);460 size_t srcFileSize = M * paddedCols * sizeof(uint16_t);
415 size_t fp4NDFileSize = M * packedCols * sizeof(uint8_t);461 size_t fp4NDFileSize = M * packedCols * sizeof(uint8_t);
462+ size_t fp4NZFileSize = M * packedCols * sizeof(uint8_t);
416 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);463 size_t e8DNFileSize = hatM * paddedCols * sizeof(uint8_t);
417 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint16_t);464 size_t maxDNFileSize = hatM * paddedCols * sizeof(uint16_t);
418 465 
@@ -423,20 +470,24 @@ void test_tquant_dn_mxfp4_fp16()
423 470 
424 uint8_t *srcHost;471 uint8_t *srcHost;
425 uint8_t *fp4NDHost;472 uint8_t *fp4NDHost;
473+ uint8_t *fp4NZHost;
426 uint8_t *e8DNHost;474 uint8_t *e8DNHost;
427 uint16_t *maxDNHost;475 uint16_t *maxDNHost;
428 uint8_t *srcDevice;476 uint8_t *srcDevice;
429 uint8_t *fp4NDDevice;477 uint8_t *fp4NDDevice;
478+ uint8_t *fp4NZDevice;
430 uint8_t *e8DNDevice;479 uint8_t *e8DNDevice;
431 uint16_t *maxDNDevice;480 uint16_t *maxDNDevice;
432 481 
433 aclrtMallocHost((void **)(&srcHost), srcFileSize);482 aclrtMallocHost((void **)(&srcHost), srcFileSize);
434 aclrtMallocHost((void **)(&fp4NDHost), fp4NDFileSize);483 aclrtMallocHost((void **)(&fp4NDHost), fp4NDFileSize);
484+ aclrtMallocHost((void **)(&fp4NZHost), fp4NZFileSize);
435 aclrtMallocHost((void **)(&e8DNHost), e8DNFileSize);485 aclrtMallocHost((void **)(&e8DNHost), e8DNFileSize);
436 aclrtMallocHost((void **)(&maxDNHost), maxDNFileSize);486 aclrtMallocHost((void **)(&maxDNHost), maxDNFileSize);
437 487 
438 aclrtMalloc((void **)&srcDevice, srcFileSize, ACL_MEM_MALLOC_HUGE_FIRST);488 aclrtMalloc((void **)&srcDevice, srcFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
439 aclrtMalloc((void **)&fp4NDDevice, fp4NDFileSize, ACL_MEM_MALLOC_HUGE_FIRST);489 aclrtMalloc((void **)&fp4NDDevice, fp4NDFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
490+ aclrtMalloc((void **)&fp4NZDevice, fp4NZFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
440 aclrtMalloc((void **)&e8DNDevice, e8DNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);491 aclrtMalloc((void **)&e8DNDevice, e8DNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
441 aclrtMalloc((void **)&maxDNDevice, maxDNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);492 aclrtMalloc((void **)&maxDNDevice, maxDNFileSize, ACL_MEM_MALLOC_HUGE_FIRST);
442 493 
@@ -445,40 +496,49 @@ void test_tquant_dn_mxfp4_fp16()
445 496 
446 const std::string goldenDir = GetGoldenDir();497 const std::string goldenDir = GetGoldenDir();
447 498 
448- TQuantDNTest::LaunchTQuantDN_MXFP4_fp16<M, N, N_pad>((uint16_t *)srcDevice, fp4NDDevice, e8DNDevice, maxDNDevice,499+ TQuantDNTest::LaunchTQuantDN_MXFP4_fp16<M, N, N_pad>((uint16_t *)srcDevice, fp4NDDevice, e8DNDevice, fp4NZDevice,
449- stream);500+ maxDNDevice, stream);
450 aclError syncRet = aclrtSynchronizeStream(stream);501 aclError syncRet = aclrtSynchronizeStream(stream);
451 ASSERT_EQ(syncRet, ACL_SUCCESS) << "MXFP4 fp16 DN sync failed: " << aclGetRecentErrMsg();502 ASSERT_EQ(syncRet, ACL_SUCCESS) << "MXFP4 fp16 DN sync failed: " << aclGetRecentErrMsg();
452 503 
453 aclrtMemcpy(fp4NDHost, fp4NDFileSize, fp4NDDevice, fp4NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);504 aclrtMemcpy(fp4NDHost, fp4NDFileSize, fp4NDDevice, fp4NDFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
505+ aclrtMemcpy(fp4NZHost, fp4NZFileSize, fp4NZDevice, fp4NZFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
454 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);506 aclrtMemcpy(e8DNHost, e8DNFileSize, e8DNDevice, e8DNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
455 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);507 aclrtMemcpy(maxDNHost, maxDNFileSize, maxDNDevice, maxDNFileSize, ACL_MEMCPY_DEVICE_TO_HOST);
456 WriteFile(goldenDir + "/output_fp4_nd.bin", fp4NDHost, fp4NDFileSize);508 WriteFile(goldenDir + "/output_fp4_nd.bin", fp4NDHost, fp4NDFileSize);
509+ WriteFile(goldenDir + "/output_fp4_nz.bin", fp4NZHost, fp4NZFileSize);
457 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);510 WriteFile(goldenDir + "/output_e8_dn.bin", e8DNHost, e8DNFileSize);
458 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);511 WriteFile(goldenDir + "/output_group_max.bin", maxDNHost, maxDNFileSize);
459 512 
460 std::vector<uint8_t> goldenFp4Nd(fp4NDFileSize);513 std::vector<uint8_t> goldenFp4Nd(fp4NDFileSize);
514+ std::vector<uint8_t> goldenFp4Nz(fp4NZFileSize);
461 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);515 std::vector<uint8_t> goldenE8Dn(e8DNFileSize);
462 std::vector<uint16_t> goldenGroupMax(maxDNFileSize / sizeof(uint16_t));516 std::vector<uint16_t> goldenGroupMax(maxDNFileSize / sizeof(uint16_t));
463 std::vector<uint8_t> outFp4Nd(fp4NDFileSize);517 std::vector<uint8_t> outFp4Nd(fp4NDFileSize);
518+ std::vector<uint8_t> outFp4Nz(fp4NZFileSize);
464 std::vector<uint8_t> outE8Dn(e8DNFileSize);519 std::vector<uint8_t> outE8Dn(e8DNFileSize);
465 std::vector<uint16_t> outGroupMax(maxDNFileSize / sizeof(uint16_t));520 std::vector<uint16_t> outGroupMax(maxDNFileSize / sizeof(uint16_t));
466 ReadFile(goldenDir + "/golden_fp4_nd.bin", fp4NDFileSize, goldenFp4Nd.data(), fp4NDFileSize);521 ReadFile(goldenDir + "/golden_fp4_nd.bin", fp4NDFileSize, goldenFp4Nd.data(), fp4NDFileSize);
522+ ReadFile(goldenDir + "/golden_fp4_nz.bin", fp4NZFileSize, goldenFp4Nz.data(), fp4NZFileSize);
467 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);523 ReadFile(goldenDir + "/golden_e8_dn.bin", e8DNFileSize, goldenE8Dn.data(), e8DNFileSize);
468 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);524 ReadFile(goldenDir + "/golden_group_max.bin", maxDNFileSize, goldenGroupMax.data(), maxDNFileSize);
469 ReadFile(goldenDir + "/output_fp4_nd.bin", fp4NDFileSize, outFp4Nd.data(), fp4NDFileSize);525 ReadFile(goldenDir + "/output_fp4_nd.bin", fp4NDFileSize, outFp4Nd.data(), fp4NDFileSize);
526+ ReadFile(goldenDir + "/output_fp4_nz.bin", fp4NZFileSize, outFp4Nz.data(), fp4NZFileSize);
470 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);527 ReadFile(goldenDir + "/output_e8_dn.bin", e8DNFileSize, outE8Dn.data(), e8DNFileSize);
471 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);528 ReadFile(goldenDir + "/output_group_max.bin", maxDNFileSize, outGroupMax.data(), maxDNFileSize);
472 ExpectGoldenMatch("MXFP4_FP16_DN", "fp4_nd", goldenFp4Nd, outFp4Nd);529 ExpectGoldenMatch("MXFP4_FP16_DN", "fp4_nd", goldenFp4Nd, outFp4Nd);
530+ ExpectGoldenMatch("MXFP4_FP16_DN", "fp4_nz", goldenFp4Nz, outFp4Nz);
473 ExpectGoldenMatch("MXFP4_FP16_DN", "e8_dn (exponents)", goldenE8Dn, outE8Dn);531 ExpectGoldenMatch("MXFP4_FP16_DN", "e8_dn (exponents)", goldenE8Dn, outE8Dn);
474 ExpectGoldenMatch("MXFP4_FP16_DN", "group_max", goldenGroupMax, outGroupMax);532 ExpectGoldenMatch("MXFP4_FP16_DN", "group_max", goldenGroupMax, outGroupMax);
475 533 
476 aclrtFree(srcDevice);534 aclrtFree(srcDevice);
477 aclrtFree(fp4NDDevice);535 aclrtFree(fp4NDDevice);
536+ aclrtFree(fp4NZDevice);
478 aclrtFree(e8DNDevice);537 aclrtFree(e8DNDevice);
479 aclrtFree(maxDNDevice);538 aclrtFree(maxDNDevice);
480 aclrtFreeHost(srcHost);539 aclrtFreeHost(srcHost);
481 aclrtFreeHost(fp4NDHost);540 aclrtFreeHost(fp4NDHost);
541+ aclrtFreeHost(fp4NZHost);
482 aclrtFreeHost(e8DNHost);542 aclrtFreeHost(e8DNHost);
483 aclrtFreeHost(maxDNHost);543 aclrtFreeHost(maxDNHost);
484 aclrtDestroyStream(stream);544 aclrtDestroyStream(stream);
@@ -19,42 +19,13 @@ using namespace pto;
19 19 
20namespace TQuantDNTest {20namespace TQuantDNTest {
21 21 
22-// Copy packed FP4 rows from the TQUANT output UB region (2D [rows, srcStride]) into22+// Full DN vector pipeline: TQuant(DN) + TMOV(ND->NZ) + TMOV<0>(DN->ZZ).
23-// the TSTORE tile UB region (2D [rows, dstStride]), honouring validPackedCols. Mirrors23+// Stores FP8 ND, E8M0 DN, per-group max, FP8 NZ, and E8M0 ZZ to GM for comparison.
24-// the proven CompactFp4PackedRows helper used by the non-DN MXFP4 kernel.24+template <typename T, int M, int N, int N_pad>
25-PTO_INTERNAL void CompactFp4PackedRows(__ubuf__ uint8_t *dstPtr, __ubuf__ uint8_t *srcPtr, uint32_t rows,
26- uint32_t validPackedCols, uint32_t srcStride, uint32_t dstStride)
27-{
28- constexpr uint32_t elementsPerRepeat = REPEAT_BYTE / sizeof(uint8_t);
29- RegTensor<uint8_t> vreg;
30- UnalignReg ureg;
31- uint16_t repeatTimes = CeilDivision(validPackedCols, elementsPerRepeat);
32- for (uint16_t row = 0; row < (uint16_t)rows; ++row) {
33- uint32_t remaining = validPackedCols;
34- __ubuf__ uint8_t *srcRow = srcPtr + row * srcStride;
35- __ubuf__ uint8_t *dstRow = dstPtr + row * dstStride;
36- for (uint16_t repeat = 0; repeat < repeatTimes; ++repeat) {
37- uint32_t cols = remaining > elementsPerRepeat ? elementsPerRepeat : remaining;
38- MaskReg preg = CreatePredicate<uint8_t>(cols);
39- uint32_t offset = repeat * elementsPerRepeat;
40- vldas(ureg, srcRow + offset);
41- vldus(vreg, ureg, srcRow + offset);
42- vsts(vreg, dstRow, offset, NORM_B8, preg);
43- remaining -= cols;
44- }
45- }
46-}
47- 
48-// Stage 1: after TQUANT (FP8 ND + E8M0 DN)
49-// Stage 2: after TQUANT + FP8 ND->NZ
50-// Stage 3: full pipeline including E8 DN->ZZ
51-template <int Stage, typename T, int M, int N, int N_pad>
52__global__ AICORE void runTQuantDN(__gm__ T __in__ *src_gm, __gm__ int8_t __out__ *fp8_nd_gm,25__global__ AICORE void runTQuantDN(__gm__ T __in__ *src_gm, __gm__ int8_t __out__ *fp8_nd_gm,
53 __gm__ uint8_t __out__ *e8_dn_gm, __gm__ int8_t __out__ *fp8_nz_gm,26 __gm__ uint8_t __out__ *e8_dn_gm, __gm__ int8_t __out__ *fp8_nz_gm,
54 __gm__ uint8_t __out__ *e8_zz_gm, __gm__ T __out__ *max_dn_gm)27 __gm__ uint8_t __out__ *e8_zz_gm, __gm__ T __out__ *max_dn_gm)
55{28{
56- static_assert(Stage >= 1 && Stage <= 3, "Stage must be 1 (quant), 2 (nz), or 3 (zz).");
57- 
58 constexpr uint32_t grpSize = 32;29 constexpr uint32_t grpSize = 32;
59 constexpr uint32_t hatM = M / grpSize;30 constexpr uint32_t hatM = M / grpSize;
60 constexpr uint32_t paddedCols = N_pad;31 constexpr uint32_t paddedCols = N_pad;
@@ -178,63 +149,48 @@ __global__ AICORE void runTQuantDN(__gm__ T __in__ *src_gm, __gm__ int8_t __out_
178 wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0);149 wait_flag(PIPE_MTE2, PIPE_V, EVENT_ID0);
179 150 
180 // Generic DN API: grp_axis=0 (groups on axis 0), single MxQuantAlg tag. The151 // Generic DN API: grp_axis=0 (groups on axis 0), single MxQuantAlg tag. The
181- // exponent is written into e8Tile; the caller reshapes it into e8DnTile via TMOV.152+ // exponent is written into e8Tile; copy it into e8DnTile so the UB tile shape
153+ // matches the GM shape exactly.
182 TQuant<0, MxQuantAlg::OcpMxFp8E4M3>(fp8Tile, srcTile, &e8Tile, &maxPerGpTile, &scalingTile);154 TQuant<0, MxQuantAlg::OcpMxFp8E4M3>(fp8Tile, srcTile, &e8Tile, &maxPerGpTile, &scalingTile);
183- 
184- // TQuant writes the exponent tile (e8Tile) in row-major [hatM, paddedCols].
185- // Copy it to e8DnTile so the UB tile shape matches the GM shape exactly.
186 TMOV(e8DnTile, e8Tile);155 TMOV(e8DnTile, e8Tile);
187 156 
188- if constexpr (Stage == 1) {157+ // Data ND->NZ (stock 2-arg TMOV) and exponent DN->ZZ (grp_axis=0 TMOV).
189- set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
190- wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
191- TSTORE(fp8NdGlobal, fp8Tile);
192- TSTORE(e8DnGlobal, e8DnTile);
193- TSTORE(maxGlobal, maxPerGpTile);
194- return;
195- }
196- 
197 TMOV(fp8TileNZ, fp8Tile);158 TMOV(fp8TileNZ, fp8Tile);
198- 159+ TMOV<0>(e8ZzTile, e8DnTile, tmpTile);
199- if constexpr (Stage == 2) {
200- set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
201- wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
202- TSTORE(fp8GlobalNZ, fp8TileNZ);
203- return;
204- }
205- 
206- TMOV(e8ZzTile, e8DnTile, tmpTile);
207 160 
208 set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);161 set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
209 wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);162 wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
210- TSTORE(e8Global, e8StoreTile);163+ TSTORE(fp8NdGlobal, fp8Tile);
164+ TSTORE(e8DnGlobal, e8DnTile);
165+ TSTORE(maxGlobal, maxPerGpTile);
211 TSTORE(fp8GlobalNZ, fp8TileNZ);166 TSTORE(fp8GlobalNZ, fp8TileNZ);
167+ TSTORE(e8Global, e8StoreTile);
212}168}
213 169 
214-template <int Stage, int M, int N, int N_pad>170+template <int M, int N, int N_pad>
215void LaunchTQuantDN(uint16_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, uint16_t *max_dn,171void LaunchTQuantDN(uint16_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, uint16_t *max_dn,
216 void *stream)172 void *stream)
217{173{
218- runTQuantDN<Stage, bfloat16_t, M, N, N_pad>174+ runTQuantDN<bfloat16_t, M, N, N_pad>
219 <<<1, nullptr, stream>>>((bfloat16_t *)src, fp8_nd, e8_dn, fp8_nz, e8_zz, (bfloat16_t *)max_dn);175 <<<1, nullptr, stream>>>((bfloat16_t *)src, fp8_nd, e8_dn, fp8_nz, e8_zz, (bfloat16_t *)max_dn);
220}176}
221 177 
222-template <int Stage, int M, int N, int N_pad>178+template <int M, int N, int N_pad>
223void LaunchTQuantDN_fp32(uint32_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz,179void LaunchTQuantDN_fp32(uint32_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz,
224 uint32_t *max_dn, void *stream)180 uint32_t *max_dn, void *stream)
225{181{
226- runTQuantDN<Stage, float, M, N, N_pad>182+ runTQuantDN<float, M, N, N_pad>
227 <<<1, nullptr, stream>>>((float *)src, fp8_nd, e8_dn, fp8_nz, e8_zz, (float *)max_dn);183 <<<1, nullptr, stream>>>((float *)src, fp8_nd, e8_dn, fp8_nz, e8_zz, (float *)max_dn);
228}184}
229 185 
230// MXFP4 (E2M1) DN kernel: quantizes src[M,N_pad] to packed FP4 plus per-group186// MXFP4 (E2M1) DN kernel: quantizes src[M,N_pad] to packed FP4 plus per-group
231-// e8m0/max tiles. Stage 3 writes FP4 into UB at byte offset (r * StaticCols + off)/2,187+// e8m0/max tiles. TQuant writes FP4 as a flat float4_e2m1x2_t tile; a uint8_t
232-// so the FP4 region is a 2D [M, packedCols] layout (packedCols = N_pad/2). A flat188+// TSTORE tile is TASSIGNed to the same UB region so TSTORE reads it in-place
233-// float4_e2m1x2_t tile is the TQUANT output; CompactFp4PackedRows bridges it to a189+// (no copy/intrinsics needed).
234-// uint8_t tile used by TSTORE (proven non-DN MXFP4 store pattern).
235template <typename T, int M, int N, int N_pad>190template <typename T, int M, int N, int N_pad>
236__global__ AICORE void runTQuantDN_MXFP4(__gm__ T __in__ *src_gm, __gm__ uint8_t __out__ *fp4_nd_gm,191__global__ AICORE void runTQuantDN_MXFP4(__gm__ T __in__ *src_gm, __gm__ uint8_t __out__ *fp4_nd_gm,
237- __gm__ uint8_t __out__ *e8_dn_gm, __gm__ T __out__ *max_dn_gm)192+ __gm__ uint8_t __out__ *e8_dn_gm, __gm__ uint8_t __out__ *fp4_nz_gm,
193+ __gm__ T __out__ *max_dn_gm)
238{194{
239 constexpr uint32_t grpSize = 32;195 constexpr uint32_t grpSize = 32;
240 constexpr uint32_t hatM = M / grpSize;196 constexpr uint32_t hatM = M / grpSize;
@@ -253,24 +209,36 @@ __global__ AICORE void runTQuantDN_MXFP4(__gm__ T __in__ *src_gm, __gm__ uint8_t
253 // TQUANT output tile: flat float4_e2m1x2_t (element = 0.5 byte -> bytes = M*packedCols).209 // TQUANT output tile: flat float4_e2m1x2_t (element = 0.5 byte -> bytes = M*packedCols).
254 using DstFP4Tile = Tile<TileType::Vec, float4_e2m1x2_t, 1, fp4FlatAligned, BLayout::RowMajor, -1, -1,210 using DstFP4Tile = Tile<TileType::Vec, float4_e2m1x2_t, 1, fp4FlatAligned, BLayout::RowMajor, -1, -1,
255 SLayout::NoneBox, 512, PadValue::Zero>;211 SLayout::NoneBox, 512, PadValue::Zero>;
256- // TSTORE tile: uint8_t [M, packedCols], same UB region bridged by CompactFp4PackedRows.212+ // 2D RowMajor view of the same packed FP4 UB data for ND->NZ.
257- // Valid extents are DYNAMIC so the tile can be runtime-sized to (M, packedCols).213+ // Use uint8_t for the view so the standard byte ND->NZ lowering is used.
258- using DstBytesTile =214+ using Fp4Tile2D = Tile<TileType::Vec, uint8_t, M, packedCols, BLayout::RowMajor, M, packedCols, SLayout::NoneBox,
259- Tile<TileType::Vec, uint8_t, M, packedCols, BLayout::RowMajor, -1, -1, SLayout::NoneBox, 512, PadValue::Zero>;215+ 512, PadValue::Zero>;
216+ // NZ tile for packed FP4: use uint8_t as the element type so the standard
217+ // byte ND->NZ lowering (32B blocks) is used. Each byte holds 2 FP4 values.
218+ constexpr uint32_t paddedRows16 = PTO_CEIL(M, FRACTAL_NZ_ROW);
219+ constexpr uint32_t virtualRow = paddedRows16 + 1;
220+ using Fp4NZTile = Tile<TileType::Vec, uint8_t, virtualRow, packedCols, BLayout::ColMajor, M, packedCols,
221+ SLayout::RowMajor, 512, PadValue::Null, CompactMode::RowPlusOne>;
222+ // TSTORE tile: uint8_t [M, packedCols] view over the same UB region that TQuant
223+ // wrote the packed FP4 data into (TASSIGN to fp4Addr). No copy needed — TSTORE
224+ // reads the bytes in-place. Valid extents are DYNAMIC for runtime sizing.
225+ using DstBytesTile = Tile<TileType::Vec, uint8_t, M, packedCols, BLayout::RowMajor, M, packedCols, SLayout::NoneBox,
226+ 512, PadValue::Zero>;
260 227 
261 constexpr uint32_t srcBytes = M * paddedCols * sizeof(T);228 constexpr uint32_t srcBytes = M * paddedCols * sizeof(T);
262 constexpr uint32_t maxBytes = hatM * paddedCols * sizeof(T);229 constexpr uint32_t maxBytes = hatM * paddedCols * sizeof(T);
263 constexpr uint32_t scalingBytes = hatM * paddedCols * sizeof(T);230 constexpr uint32_t scalingBytes = hatM * paddedCols * sizeof(T);
264 constexpr uint32_t e8Bytes = hatM * paddedCols;231 constexpr uint32_t e8Bytes = hatM * paddedCols;
265 constexpr uint32_t fp4Bytes = M * packedCols;232 constexpr uint32_t fp4Bytes = M * packedCols;
233+ constexpr uint32_t fp4NZBytes = virtualRow * packedCols;
266 234 
267 constexpr uint32_t srcAddr = 0x0;235 constexpr uint32_t srcAddr = 0x0;
268 constexpr uint32_t maxAddr = PTO_CEIL(srcAddr + srcBytes, 0x20);236 constexpr uint32_t maxAddr = PTO_CEIL(srcAddr + srcBytes, 0x20);
269 constexpr uint32_t scalingAddr = PTO_CEIL(maxAddr + maxBytes, 0x20);237 constexpr uint32_t scalingAddr = PTO_CEIL(maxAddr + maxBytes, 0x20);
270 constexpr uint32_t e8Addr = PTO_CEIL(scalingAddr + scalingBytes, 0x20);238 constexpr uint32_t e8Addr = PTO_CEIL(scalingAddr + scalingBytes, 0x20);
271 constexpr uint32_t fp4Addr = PTO_CEIL(e8Addr + e8Bytes, 0x20);239 constexpr uint32_t fp4Addr = PTO_CEIL(e8Addr + e8Bytes, 0x20);
272- constexpr uint32_t fp4StoreAddr = PTO_CEIL(fp4Addr + fp4Bytes, 0x20);240+ constexpr uint32_t fp4NZAddr = PTO_CEIL(fp4Addr + fp4Bytes, 0x20);
273- constexpr uint32_t layoutEnd = PTO_CEIL(fp4StoreAddr + fp4Bytes, 0x100);241+ constexpr uint32_t layoutEnd = PTO_CEIL(fp4NZAddr + fp4NZBytes, 0x100);
274 static_assert(layoutEnd <= 0x40000, "MXFP4 DN UB layout exceeds 256 KB.");242 static_assert(layoutEnd <= 0x40000, "MXFP4 DN UB layout exceeds 256 KB.");
275 243 
276 SrcTile srcTile(M, paddedCols);244 SrcTile srcTile(M, paddedCols);
@@ -278,22 +246,29 @@ __global__ AICORE void runTQuantDN_MXFP4(__gm__ T __in__ *src_gm, __gm__ uint8_t
278 ScalingTile scalingTile(hatM, paddedCols);246 ScalingTile scalingTile(hatM, paddedCols);
279 E8Tile e8Tile(hatM, paddedCols);247 E8Tile e8Tile(hatM, paddedCols);
280 DstFP4Tile fp4Tile;248 DstFP4Tile fp4Tile;
281- DstBytesTile fp4BytesTile(M, packedCols);249+ Fp4Tile2D fp4Tile2D;
250+ Fp4NZTile fp4NZTile;
251+ DstBytesTile fp4BytesTile;
282 252 
283 TASSIGN(srcTile, srcAddr);253 TASSIGN(srcTile, srcAddr);
284 TASSIGN(maxTile, maxAddr);254 TASSIGN(maxTile, maxAddr);
285 TASSIGN(scalingTile, scalingAddr);255 TASSIGN(scalingTile, scalingAddr);
286 TASSIGN(e8Tile, e8Addr);256 TASSIGN(e8Tile, e8Addr);
287 TASSIGN(fp4Tile, fp4Addr);257 TASSIGN(fp4Tile, fp4Addr);
288- TASSIGN(fp4BytesTile, fp4StoreAddr);258+ TASSIGN(fp4Tile2D, fp4Addr);
259+ TASSIGN(fp4NZTile, fp4NZAddr);
260+ TASSIGN(fp4BytesTile, fp4Addr);
289 261 
290 using SrcGlobal = GlobalTensor<T, Shape<1, 1, 1, M, N_pad>, pto::Stride<1, 1, 1, N_pad, 1>>;262 using SrcGlobal = GlobalTensor<T, Shape<1, 1, 1, M, N_pad>, pto::Stride<1, 1, 1, N_pad, 1>>;
291 using Fp4Global = GlobalTensor<uint8_t, Shape<1, 1, 1, M, packedCols>, pto::Stride<1, 1, 1, packedCols, 1>>;263 using Fp4Global = GlobalTensor<uint8_t, Shape<1, 1, 1, M, packedCols>, pto::Stride<1, 1, 1, packedCols, 1>>;
264+ using Fp4GlobalNZ = GlobalTensor<uint8_t, TileShape2D<uint8_t, M, packedCols, Layout::NZ>,
265+ BaseShape2D<uint8_t, M, packedCols, Layout::NZ>, Layout::NZ>;
292 using MaxGlobal = GlobalTensor<T, Shape<1, 1, 1, hatM, paddedCols>, pto::Stride<1, 1, 1, paddedCols, 1>>;266 using MaxGlobal = GlobalTensor<T, Shape<1, 1, 1, hatM, paddedCols>, pto::Stride<1, 1, 1, paddedCols, 1>>;
293 using E8Global = GlobalTensor<uint8_t, Shape<1, 1, 1, hatM, paddedCols>, pto::Stride<1, 1, 1, paddedCols, 1>>;267 using E8Global = GlobalTensor<uint8_t, Shape<1, 1, 1, hatM, paddedCols>, pto::Stride<1, 1, 1, paddedCols, 1>>;
294 268 
295 SrcGlobal srcGlobal(src_gm);269 SrcGlobal srcGlobal(src_gm);
296 Fp4Global fp4Global(fp4_nd_gm);270 Fp4Global fp4Global(fp4_nd_gm);
271+ Fp4GlobalNZ fp4GlobalNZ(fp4_nz_gm);
297 MaxGlobal maxGlobal(max_dn_gm);272 MaxGlobal maxGlobal(max_dn_gm);
298 E8Global e8Global(e8_dn_gm);273 E8Global e8Global(e8_dn_gm);
299 274 
@@ -305,60 +280,56 @@ __global__ AICORE void runTQuantDN_MXFP4(__gm__ T __in__ *src_gm, __gm__ uint8_t
305 // into e8Tile; the kernel TSTOREs e8Tile directly (shape already matches GM).280 // into e8Tile; the kernel TSTOREs e8Tile directly (shape already matches GM).
306 TQuant<0, MxQuantAlg::OcpMxFp4E2M1>(fp4Tile, srcTile, &e8Tile, &maxTile, &scalingTile);281 TQuant<0, MxQuantAlg::OcpMxFp4E2M1>(fp4Tile, srcTile, &e8Tile, &maxTile, &scalingTile);
307 282 
308- __VEC_SCOPE__283+ // Packed FP4 ND->NZ: source is RowMajor [M, packedCols] of float4_e2m1x2_t.
309- {284+ TMOV(fp4NZTile, fp4Tile2D);
310- // Stage 3 wrote FP4 densely at [r, packedCols] (srcStride == dstStride == packedCols),
311- // so the compact copy is a plain row-wise move into the TSTORE tile's UB region.
312- mem_bar(VST_VLD);
313- CompactFp4PackedRows((__ubuf__ uint8_t *)fp4BytesTile.data(), (__ubuf__ uint8_t *)fp4Tile.data(), M, packedCols,
314- packedCols, packedCols);
315- mem_bar(VST_VST);
316- }
317 285 
318 set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);286 set_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
319 wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);287 wait_flag(PIPE_V, PIPE_MTE3, EVENT_ID0);
320 TSTORE(fp4Global, fp4BytesTile);288 TSTORE(fp4Global, fp4BytesTile);
289+ TSTORE(fp4GlobalNZ, fp4NZTile);
321 TSTORE(maxGlobal, maxTile);290 TSTORE(maxGlobal, maxTile);
322 TSTORE(e8Global, e8Tile);291 TSTORE(e8Global, e8Tile);
323}292}
324 293 
325template <int M, int N, int N_pad>294template <int M, int N, int N_pad>
326-void LaunchTQuantDN_MXFP4_bf16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint16_t *max_dn, void *stream)295+void LaunchTQuantDN_MXFP4_bf16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint8_t *fp4_nz, uint16_t *max_dn,
296+ void *stream)
327{297{
328 runTQuantDN_MXFP4<bfloat16_t, M, N, N_pad>298 runTQuantDN_MXFP4<bfloat16_t, M, N, N_pad>
329- <<<1, nullptr, stream>>>((bfloat16_t *)src, fp4_nd, e8_dn, (bfloat16_t *)max_dn);299+ <<<1, nullptr, stream>>>((bfloat16_t *)src, fp4_nd, e8_dn, fp4_nz, (bfloat16_t *)max_dn);
330}300}
331 301 
332template <int M, int N, int N_pad>302template <int M, int N, int N_pad>
333-void LaunchTQuantDN_MXFP4_fp16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint16_t *max_dn, void *stream)303+void LaunchTQuantDN_MXFP4_fp16(uint16_t *src, uint8_t *fp4_nd, uint8_t *e8_dn, uint8_t *fp4_nz, uint16_t *max_dn,
304+ void *stream)
334{305{
335- runTQuantDN_MXFP4<half, M, N, N_pad><<<1, nullptr, stream>>>((half *)src, fp4_nd, e8_dn, (half *)max_dn);306+ runTQuantDN_MXFP4<half, M, N, N_pad><<<1, nullptr, stream>>>((half *)src, fp4_nd, e8_dn, fp4_nz, (half *)max_dn);
336}307}
337 308 
338-#define INSTANTIATE_TQUANT_DN_STAGE(S, M, N, NP) \309+#define INSTANTIATE_TQUANT_DN(M, N, NP) \
339- template void LaunchTQuantDN<S, M, N, NP>(uint16_t *, int8_t *, uint8_t *, int8_t *, uint8_t *, uint16_t *, void *)310+ template void LaunchTQuantDN<M, N, NP>(uint16_t *, int8_t *, uint8_t *, int8_t *, uint8_t *, uint16_t *, void *)
340 311 
341-#define INSTANTIATE_TQUANT_DN_STAGE_FP32(S, M, N, NP) \312+#define INSTANTIATE_TQUANT_DN_FP32(M, N, NP) \
342- template void LaunchTQuantDN_fp32<S, M, N, NP>(uint32_t *, int8_t *, uint8_t *, int8_t *, uint8_t *, uint32_t *, \313+ template void LaunchTQuantDN_fp32<M, N, NP>(uint32_t *, int8_t *, uint8_t *, int8_t *, uint8_t *, uint32_t *, \
343- void *)314+ void *)
344 315 
345-INSTANTIATE_TQUANT_DN_STAGE(1, 128, 128, 128);316+INSTANTIATE_TQUANT_DN(128, 128, 128);
346-INSTANTIATE_TQUANT_DN_STAGE(1, 64, 128, 128);317+INSTANTIATE_TQUANT_DN(64, 128, 128);
347-INSTANTIATE_TQUANT_DN_STAGE(1, 64, 256, 256);318+INSTANTIATE_TQUANT_DN(64, 256, 256);
348-INSTANTIATE_TQUANT_DN_STAGE(1, 128, 256, 256);319+INSTANTIATE_TQUANT_DN(128, 256, 256);
349-INSTANTIATE_TQUANT_DN_STAGE(1, 64, 64, 64);320+INSTANTIATE_TQUANT_DN(64, 64, 64);
350-INSTANTIATE_TQUANT_DN_STAGE(1, 128, 64, 64);321+INSTANTIATE_TQUANT_DN(128, 64, 64);
351-INSTANTIATE_TQUANT_DN_STAGE(1, 256, 64, 64);322+INSTANTIATE_TQUANT_DN(256, 64, 64);
352-INSTANTIATE_TQUANT_DN_STAGE(1, 256, 128, 128);323+INSTANTIATE_TQUANT_DN(256, 128, 128);
353-INSTANTIATE_TQUANT_DN_STAGE_FP32(1, 64, 128, 128);324+INSTANTIATE_TQUANT_DN_FP32(64, 128, 128);
354-INSTANTIATE_TQUANT_DN_STAGE_FP32(1, 128, 128, 128);325+INSTANTIATE_TQUANT_DN_FP32(128, 128, 128);
355-INSTANTIATE_TQUANT_DN_STAGE_FP32(1, 64, 256, 256);326+INSTANTIATE_TQUANT_DN_FP32(64, 256, 256);
356 327 
357#define INSTANTIATE_TQUANT_DN_MXFP4_BF16(M, N, NP) \328#define INSTANTIATE_TQUANT_DN_MXFP4_BF16(M, N, NP) \
358- template void LaunchTQuantDN_MXFP4_bf16<M, N, NP>(uint16_t *, uint8_t *, uint8_t *, uint16_t *, void *)329+ template void LaunchTQuantDN_MXFP4_bf16<M, N, NP>(uint16_t *, uint8_t *, uint8_t *, uint8_t *, uint16_t *, void *)
359 330 
360#define INSTANTIATE_TQUANT_DN_MXFP4_FP16(M, N, NP) \331#define INSTANTIATE_TQUANT_DN_MXFP4_FP16(M, N, NP) \
361- template void LaunchTQuantDN_MXFP4_fp16<M, N, NP>(uint16_t *, uint8_t *, uint8_t *, uint16_t *, void *)332+ template void LaunchTQuantDN_MXFP4_fp16<M, N, NP>(uint16_t *, uint8_t *, uint8_t *, uint8_t *, uint16_t *, void *)
362 333 
363INSTANTIATE_TQUANT_DN_MXFP4_BF16(64, 128, 128);334INSTANTIATE_TQUANT_DN_MXFP4_BF16(64, 128, 128);
364INSTANTIATE_TQUANT_DN_MXFP4_BF16(128, 128, 128);335INSTANTIATE_TQUANT_DN_MXFP4_BF16(128, 128, 128);
@@ -367,6 +338,6 @@ INSTANTIATE_TQUANT_DN_MXFP4_FP16(64, 128, 128);
367INSTANTIATE_TQUANT_DN_MXFP4_FP16(128, 128, 128);338INSTANTIATE_TQUANT_DN_MXFP4_FP16(128, 128, 128);
368INSTANTIATE_TQUANT_DN_MXFP4_FP16(64, 256, 256);339INSTANTIATE_TQUANT_DN_MXFP4_FP16(64, 256, 256);
369 340 
370-#undef INSTANTIATE_TQUANT_DN_STAGE341+#undef INSTANTIATE_TQUANT_DN
371 342 
372} // namespace TQuantDNTest343} // namespace TQuantDNTest