已合并
[A5]: Adding support for TMOV DN to ZZ #1196
omarzohir创建于 6月26日
[A5]: Adding support for TMOV DN to ZZ #1196
已合并
共 7 个文件变更+518-168
| @@ -0,0 +1,159 @@ | |||
| 1 | +# TQUANT DN — Axis-0 Grouped Quantization and DN→ZZ | ||
| 2 | + | ||
| 3 | +## Tile Operation Diagram | ||
| 4 | + | ||
| 5 | + | ||
| 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 | + | ||
| 1278 | template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, typename... WaitEvents> | 1290 | template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, typename... WaitEvents> |
| 1279 | PTO_INST RecordEvent TMOV(DstTileData &dst, SrcTileData &src, WaitEvents &...events) | 1291 | PTO_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 | + | ||
| 380 | template <typename WorkT, typename SrcTileData> | 460 | template <typename WorkT, typename SrcTileData> |
| 381 | PTO_INTERNAL void TMovNd2NzLoop(__ubuf__ WorkT *srcPtr, __ubuf__ WorkT *dstPtr, uint16_t repeatTimes, | 461 | PTO_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> |
| 615 | PTO_INTERNAL void TMOV_IMPL(DstTileData &dst, SrcTileData &src, TmpTileData &tmp) | 702 | PTO_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 | ||
| 622 | template <typename DstTileData, typename SrcTileData, ReluPreMode reluMode, STPhase Phase = STPhase::Unspecified> | 714 | template <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 | ||
| 83 | if(DEBUG_MODE) | 82 | if(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 | ||
| 13 | import math | 13 | import math |
| 14 | import os | 14 | import os |
| 15 | +from dataclasses import dataclass | ||
| 16 | +from typing import Optional | ||
| 15 | 17 | ||
| 16 | import numpy as np | 18 | import numpy as np |
| 17 | from ml_dtypes import bfloat16, float4_e2m1fn | 19 | from ml_dtypes import bfloat16, float4_e2m1fn |
| @@ -73,21 +75,14 @@ def fp32_to_e4m3(x): | |||
| 73 | 75 | ||
| 74 | 76 | ||
| 75 | def nd2nz_mxfp8(data_fp8, m, n): | 77 | def nd2nz_mxfp8(data_fp8, m, n): |
| 76 | - padded_rows16 = ((m + 15) // 16) * 16 | 78 | + # Stock ND->NZ for 1-byte data: [M,N] -> [n_groups, padded_m, 32]. No virtual_row+1 |
| 77 | - virtual_row = padded_rows16 + 1 | 79 | + # (the +1 is a UB-internal stride; the GM NZ layout is plain [n_groups, padded_m, 32]). |
| 78 | - padded_cols = ((n + 31) // 32) * 32 | 80 | + padded_m = ((m + 15) // 16) * 16 |
| 79 | - n_col_groups = padded_cols // 32 | 81 | + 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_fp8 | 83 | + 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 | ||
| 93 | def pack_e8_dn(e8m0, hat_m, n, padded_cols): | 88 | def 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 | ||
| 102 | def dn2zz_e8m0(e8m0_dn, hat_m, n): | 97 | def 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() |
| 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 | ||
| 107 | def quant_bf16_to_mxfp8_dn(src_bf16_fp32, m, n_pad): | 105 | def 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.0 | 308 | 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 | + | ||
| 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 | ||
| 325 | def gen_golden_data(case_name, m, n): | 348 | def 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 | ||
| 340 | def gen_golden_data_fp32(case_name, m, n): | 371 | def gen_golden_data_fp32(case_name, m, n): |
| 341 | n_pad = n | 372 | 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 | ||
| 366 | def gen_golden_data_mxfp4_bf16(case_name, m, n): | 407 | def 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 | ||
| 381 | def _gen_src_fp16_safe(m, n_pad): | 430 | def _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 MaxTile | 455 | # 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 | ||
| 413 | if __name__ == "__main__": | 470 | if __name__ == "__main__": |
| @@ -17,19 +17,21 @@ using namespace PtoTestCommon; | |||
| 17 | 17 | ||
| 18 | namespace TQuantDNTest { | 18 | namespace TQuantDNTest { |
| 19 | 19 | ||
| 20 | -template <int Stage, int M, int N, int N_pad> | 20 | +template <int M, int N, int N_pad> |
| 21 | void LaunchTQuantDN(uint16_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, uint16_t *max_dn, | 21 | void 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> |
| 25 | void LaunchTQuantDN_fp32(uint32_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, | 25 | void 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 | ||
| 28 | template <int M, int N, int N_pad> | 28 | template <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 | ||
| 31 | template <int M, int N, int N_pad> | 32 | template <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 TQuantDNTest | 36 | } // 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 max | 118 | + // 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 bytes | 342 | // e8_dn : hatM * paddedCols bytes |
| 311 | // max_dn : hatM * paddedCols * sizeof(b16) bytes | 343 | // max_dn : hatM * paddedCols * sizeof(b16) bytes |
| 312 | template <int M, int N, int N_pad> | 344 | template <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 | ||
| 20 | namespace TQuantDNTest { | 20 | namespace TQuantDNTest { |
| 21 | 21 | ||
| 22 | -// Copy packed FP4 rows from the TQUANT output UB region (2D [rows, srcStride]) into | 22 | +// Full DN vector pipeline: TQuant(DN) + TMOV(ND->NZ) + TMOV<0>(DN->ZZ). |
| 23 | -// the TSTORE tile UB region (2D [rows, dstStride]), honouring validPackedCols. Mirrors | 23 | +// 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. The | 151 | // 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> |
| 215 | void LaunchTQuantDN(uint16_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, uint16_t *max_dn, | 171 | void 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> |
| 223 | void LaunchTQuantDN_fp32(uint32_t *src, int8_t *fp8_nd, uint8_t *e8_dn, int8_t *fp8_nz, uint8_t *e8_zz, | 179 | void 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-group | 186 | // 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 flat | 188 | +// 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 a | 189 | +// (no copy/intrinsics needed). |
| 234 | -// uint8_t tile used by TSTORE (proven non-DN MXFP4 store pattern). | ||
| 235 | template <typename T, int M, int N, int N_pad> | 190 | template <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 | ||
| 325 | template <int M, int N, int N_pad> | 294 | template <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 | ||
| 332 | template <int M, int N, int N_pad> | 302 | template <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 | 328 | ||
| 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 | 331 | ||
| 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 | ||
| 363 | INSTANTIATE_TQUANT_DN_MXFP4_BF16(64, 128, 128); | 334 | INSTANTIATE_TQUANT_DN_MXFP4_BF16(64, 128, 128); |
| 364 | INSTANTIATE_TQUANT_DN_MXFP4_BF16(128, 128, 128); | 335 | INSTANTIATE_TQUANT_DN_MXFP4_BF16(128, 128, 128); |
| @@ -367,6 +338,6 @@ INSTANTIATE_TQUANT_DN_MXFP4_FP16(64, 128, 128); | |||
| 367 | INSTANTIATE_TQUANT_DN_MXFP4_FP16(128, 128, 128); | 338 | INSTANTIATE_TQUANT_DN_MXFP4_FP16(128, 128, 128); |
| 368 | INSTANTIATE_TQUANT_DN_MXFP4_FP16(64, 256, 256); | 339 | INSTANTIATE_TQUANT_DN_MXFP4_FP16(64, 256, 256); |
| 369 | 340 | ||
| 370 | -#undef INSTANTIATE_TQUANT_DN_STAGE | 341 | +#undef INSTANTIATE_TQUANT_DN |
| 371 | 342 | ||
| 372 | } // namespace TQuantDNTest | 343 | } // namespace TQuantDNTest |
🟡 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) 匹配。