Skip to content

feat(hir,analysis,tir): support block-scaled fp8 gemm on sm90 - #212

Open
zhen8838 wants to merge 3 commits into
tile-ai:mainfrom
zhen8838:feat/sm90-fp8-block-scaled-gemm
Open

zhen8838 wants to merge 3 commits into
tile-ai:mainfrom
zhen8838:feat/sm90-fp8-block-scaled-gemm

Conversation

@zhen8838

@zhen8838 zhen8838 commented Oct 5, 2026

Copy link
Copy Markdown
Collaborator

Why

What

  • MatMul types, costs and evaluates a product of two fp8e4m3 operands in
    f32 (matmul_result_dtype). Every other dtype keeps its own result dtype.
  • DimAdd, DimSub, DimMul, DimFloorDiv, DimMod, DimMin and DimMax
    count as integer work in the compute cost.
  • New T.cuda.sm90.WgmmaFp8(n=...): wgmma.mma_async 64 x n x 32 for e4m3,
    f32 accumulator fragment as Wgmma, A (64, 32) and B (32, n) K-major in
    shared memory, in swizzled rows of 16 to 128 bytes. It
    joins the MMA family, so schedule candidates offers it and finalize lowers
    it. The schedule evaluator sums FP8 operands in f32.
  • schedule facts lists each capability once (two atoms share
    wgmma.mma_async).
  • Tests: MatMul typeinfer/cost/eval cases, a dimension-arithmetic slice
    analysis case, an FP8 K-major WGMMA fixture with its TIR golden, facts
    golden, smem/rmem budgets, and an exact-f32 evaluation check. The plain
    candidates golden gains the two WgmmaFp8 refusal lines.

Contract

  • docs/spec/hir.md: MatMul returns f32 for fp8e4m3 operands. This changes
    the result dtype of any program that already multiplied fp8 tensors; before,
    such a program could not combine the product with an f32 value.
  • docs/spec/tir.md and docs/spec/code-organization.md declare
    WgmmaFp8 and its file.
  • No change for bf16, f16 or f32 programs.

Risk

  • Only e4m3 x e4m3 is declared; e5m2 and mixed FP8 operands are not.
  • The declaration reads A from shared memory only (no register-A form), and
    B must be K-major; an (N, K) row-major weight is written as its (K, N) view.
  • The CUDA build tests in this environment fail at the base commit too (no
    CUTLASS checkout); the schedule, analysis, evaluator and MatMul suites pass
    (464 tests).

tf.matmul of two fp8e4m3 operands returned fp8e4m3, so a block-scaled
FP8 GEMM could not multiply its product by an f32 scale. No MMA sums in
fp8e4m3: MatMul now types, costs and evaluates fp8e4m3 products in f32.
A slice bound written with dimension arithmetic, such as
scale[:, gi * G:gi * G + G], stopped analyze with 'DimMul: no cost
evaluator registered'. DimAdd, DimSub, DimMul, DimFloorDiv, DimMod, DimMin
and DimMax now count as integer work.
SM90 had only the bf16/f16 Wgmma atom, so tiled_mma refused fp8e4m3
operands and schedule candidates offered no FP8 instruction.
T.cuda.sm90.WgmmaFp8 declares the 64 x n x 32 e4m3 MMA with an f32
accumulator and K-major A and B in swizzled shared memory. The schedule
evaluator sums its operands in f32, and facts lists each capability once.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

feat(hir,analysis,tir): support block-scaled fp8 gemm on sm90

1 participant