Skip to content

Commit 5c70720

Browse files
committed
[water] Add ScaledMMAOp
Implement ScaledMMA Op that operates on MXFP data types like MXFP4. Signed-off-by: Tim Gymnich <tim@gymni.ch>
1 parent a644ccb commit 5c70720

10 files changed

Lines changed: 1027 additions & 15 deletions

File tree

lit_tests/kernel/wave/mlir_converter.py

Lines changed: 137 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@
2222
wave_compile,
2323
)
2424
from wave_lang.kernel.wave.type_inference import infer_types
25-
from wave_lang.kernel.wave.constraints import Constraint, MMAType
25+
from wave_lang.kernel.wave.constraints import Constraint, MMAType, ScaledMMAType
2626
from wave_lang.kernel.wave.mlir_converter.mlir_converter import (
2727
format_diagnostics,
2828
PersistentEmitter,
@@ -1866,3 +1866,139 @@ def mlir_converter_attention_pre_infer_types():
18661866

18671867
# Iterate result types match init_arg types.
18681868
# CHECK: -> (!wave.tensor<[@B, @M] of f32, <register>>, !wave.tensor<[@B, @M] of f32, <register>>, !wave.tensor<[@B, @N, @M] of f32, <register>>)
1869+
1870+
1871+
@run_test
1872+
def test_mxfp4_scaled_mma_256x256x256():
1873+
# Input sizes
1874+
M = tkl.sym.M
1875+
N = tkl.sym.N
1876+
K = tkl.sym.K
1877+
# Workgroup tile sizes
1878+
BLOCK_M = tkl.sym.BLOCK_M
1879+
BLOCK_N = tkl.sym.BLOCK_N
1880+
BLOCK_K = tkl.sym.BLOCK_K
1881+
# Address space (for GPU, shared(1) or global(0))
1882+
ADDRESS_SPACE = tkl.sym.ADDRESS_SPACE
1883+
1884+
mfma_variant = ScaledMMAType.F32_16x16x128_F8F6F4
1885+
1886+
# Expose user-constraints
1887+
constraints: list[tkw.Constraint] = [tkw.WorkgroupConstraint(M, BLOCK_M, 0)]
1888+
constraints += [tkw.WorkgroupConstraint(N, BLOCK_N, 1)]
1889+
constraints += [tkw.TilingConstraint(K, BLOCK_K)]
1890+
constraints += [tkw.WaveConstraint(M, BLOCK_M / 2)]
1891+
constraints += [tkw.WaveConstraint(N, BLOCK_N / 2)]
1892+
1893+
constraints += [tkw.HardwareConstraint(threads_per_wave=64, mma_type=mfma_variant)]
1894+
1895+
@tkw.wave(constraints)
1896+
def scaled_gemm(
1897+
a: tkl.Memory[M, K / 2, ADDRESS_SPACE, tkl.i8],
1898+
a_scale: tkl.Memory[M, K / 32, ADDRESS_SPACE, tkl.i8],
1899+
b: tkl.Memory[N, K / 2, ADDRESS_SPACE, tkl.i8],
1900+
b_scale: tkl.Memory[N, K / 32, ADDRESS_SPACE, tkl.i8],
1901+
c: tkl.Memory[M, N, GLOBAL_ADDRESS_SPACE, tkl.f32],
1902+
):
1903+
c_reg = tkl.Register[M, N, tkl.f32](0.0)
1904+
1905+
@tkw.iterate(K, init_args=[c_reg])
1906+
def repeat(acc: tkl.Register[M, N, tkl.f32]) -> tkl.Register[M, N, tkl.f32]:
1907+
a_reg = tkw.read(a)
1908+
a_reg = tkw.bitcast(a_reg, tkl.f4e2m1fn)
1909+
a_scale_reg = tkw.read(a_scale)
1910+
a_scale_reg = tkw.bitcast(a_scale_reg, tkl.f8e8m0fnu)
1911+
b_reg = tkw.read(b)
1912+
b_reg = tkw.bitcast(b_reg, tkl.f4e2m1fn)
1913+
b_scale_reg = tkw.read(b_scale)
1914+
b_scale_reg = tkw.bitcast(b_scale_reg, tkl.f8e8m0fnu)
1915+
acc = tkw.scaled_mma(a_reg, a_scale_reg, b_reg, b_scale_reg, acc)
1916+
return acc
1917+
1918+
tkw.write(repeat, c)
1919+
1920+
subs = {
1921+
ADDRESS_SPACE: SHARED_ADDRESS_SPACE,
1922+
BLOCK_M: 32,
1923+
BLOCK_N: 32,
1924+
BLOCK_K: 256,
1925+
M: 16384,
1926+
N: 16384,
1927+
K: 16384,
1928+
}
1929+
options = WaveCompileOptions(
1930+
subs=subs,
1931+
compile_to_mlir=True,
1932+
location_capture_config=LocationCaptureConfig(level=LocationCaptureLevel.NONE),
1933+
enforce_locations=False,
1934+
)
1935+
options = set_default_run_config(options)
1936+
1937+
compiled_kernel = wave_compile(options, scaled_gemm)
1938+
trace = compiled_kernel.get_compiled_graph()
1939+
1940+
mlir_output, diagnostics, _ = emitter.emit_wave_dialect(
1941+
trace, scaled_gemm.constraints, options
1942+
)
1943+
1944+
if diagnostics:
1945+
print(format_diagnostics(diagnostics, use_color=False), file=sys.stderr)
1946+
assert (
1947+
len(diagnostics) == 0
1948+
), "dialect emission should create valid IR, therefore diagnostics should be empty"
1949+
1950+
print(mlir_output)
1951+
1952+
# CHECK-LABEL: test_mxfp4_scaled_mma_256x256x256
1953+
# CHECK: module
1954+
# CHECK-NEXT: func.func @kernel(
1955+
# CHECK-SAME: %[[A:.*]]: !wave.tensor<[@M, @K2] of i8, <global>>
1956+
# CHECK-SAME: %[[A_SCALE:.*]]: !wave.tensor<[@M, @K32] of i8, <global>>
1957+
# CHECK-SAME: %[[B:.*]]: !wave.tensor<[@N, @K2] of i8, <global>>
1958+
# CHECK-SAME: %[[B_SCALE:.*]]: !wave.tensor<[@N, @K32] of i8, <global>>
1959+
# CHECK-SAME: %[[C:.*]]: !wave.tensor<[@M, @N] of f32, <global>>
1960+
# CHECK-SAME: wave.constraints =
1961+
# CHECK-SAME: #wave.workgroup_constraint<dim = <"M">, tile_size = <[#wave.symbol<"BLOCK_M">] -> (BLOCK_M)>, workgroup_dim = <x>>
1962+
# CHECK-SAME: #wave.workgroup_constraint<dim = <"N">, tile_size = <[#wave.symbol<"BLOCK_N">] -> (BLOCK_N)>, workgroup_dim = <y>>
1963+
# CHECK-SAME: #wave.tiling_constraint<dim = <"K">, tile_size = <[#wave.symbol<"BLOCK_K">] -> (BLOCK_K)>>
1964+
# CHECK-SAME: #wave.wave_constraint<dim = <"M">, tile_size = <[#wave.symbol<"BLOCK_M">] -> (BLOCK_M floordiv 2)>>
1965+
# CHECK-SAME: #wave.wave_constraint<dim = <"N">, tile_size = <[#wave.symbol<"BLOCK_N">] -> (BLOCK_N floordiv 2)>>
1966+
# CHECK-SAME: #wave.hardware_constraint<threads_per_wave = 64, waves_per_block = [2, 2, 1], mma_type = <f32_16x16x128_f8f6f4>>
1967+
# CHECK-SAME: #wave.hyperparameters<{BLOCK_K = 256 : i64, BLOCK_M = 32 : i64, BLOCK_N = 32 : i64, K = 16384 : i64, K2 = #wave.expr_list<[#wave.symbol<"K">] -> (K floordiv 2)>, K32 = #wave.expr_list<[#wave.symbol<"K">] -> (K floordiv 32)>, M = 16384 : i64, N = 16384 : i64}>
1968+
#
1969+
# CHECK: %[[CST:.*]] = arith.constant 0.000000e+00 : f32
1970+
# CHECK-NEXT: %[[REG:.*]] = wave.register %[[CST]]
1971+
#
1972+
# CHECK: %[[ITERATE:.*]] = wave.iterate @K iter_args(%[[REG]]) {
1973+
# CHECK-NEXT: ^{{.*}}(%[[ACC:.*]]: !wave.tensor<[@M, @N] of f32, <register>>):
1974+
#
1975+
# Global reads promoted through shared memory.
1976+
#
1977+
# CHECK: wave.read %[[A]]
1978+
# CHECK: wave.write {{.*}} !wave.tensor<[@M, @K2] of i8, <shared>>
1979+
# CHECK: wave.read %[[A_SCALE]]
1980+
# CHECK: wave.write {{.*}} !wave.tensor<[@M, @K32] of i8, <shared>>
1981+
# CHECK: wave.read %[[B]]
1982+
# CHECK: wave.write {{.*}} !wave.tensor<[@N, @K2] of i8, <shared>>
1983+
# CHECK: wave.read %[[B_SCALE]]
1984+
# CHECK: wave.write {{.*}} !wave.tensor<[@N, @K32] of i8, <shared>>
1985+
#
1986+
# Bitcasts from i8 to scaled MMA operand types.
1987+
#
1988+
# CHECK: wave.bitcast {{.*}} : !wave.tensor<[@M, @K2] of i8, <register>> to !wave.tensor<[@M, @K] of f4E2M1FN, <register>>
1989+
# CHECK: wave.bitcast {{.*}} : !wave.tensor<[@M, @K32] of i8, <register>> to !wave.tensor<[@M, @K32] of f8E8M0FNU, <register>>
1990+
# CHECK: wave.bitcast {{.*}} : !wave.tensor<[@N, @K2] of i8, <register>> to !wave.tensor<[@N, @K] of f4E2M1FN, <register>>
1991+
# CHECK: wave.bitcast {{.*}} : !wave.tensor<[@N, @K32] of i8, <register>> to !wave.tensor<[@N, @K32] of f8E8M0FNU, <register>>
1992+
#
1993+
# Two scaled MMAs (BLOCK_K=256 with 128 per intrinsic).
1994+
#
1995+
# CHECK: %[[SMMA0:.*]] = wave.scaled_mma {{.*}}, {{.*}}, {{.*}}, {{.*}}, %[[ACC]]
1996+
# CHECK-SAME: #wave.mma_kind<f32_16x16x128_f8f6f4>
1997+
# CHECK: %[[SMMA1:.*]] = wave.scaled_mma {{.*}}, {{.*}}, {{.*}}, {{.*}}, %[[SMMA0]]
1998+
# CHECK-SAME: #wave.mma_kind<f32_16x16x128_f8f6f4>
1999+
# CHECK: wave.yield %[[SMMA1]] : !wave.tensor<[@M, @N] of f32, <register>>
2000+
# CHECK-NEXT: }
2001+
#
2002+
# Results written back to the output tensor.
2003+
#
2004+
# CHECK: wave.write {{.*}}, %[[C]]

water/include/water/Dialect/Wave/IR/WaveOps.td

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -244,6 +244,51 @@ def MmaOp : WaveOp<"mma",
244244
}];
245245
}
246246

247+
def ScaledMmaOp : WaveOp<"scaled_mma",
248+
[DeclareOpInterfaceMethods<WaveInferTypeOpInterface>,
249+
DeclareOpInterfaceMethods<WaveElementsPerThreadOpInterface>,
250+
DeclareOpInterfaceMethods<WaveInferIndexExprsOpInterface,
251+
["initializeIndexExprsForward", "initializeIndexExprsBackward",
252+
"getIndexExprValuesAndDescriptions"]>]>,
253+
WaveArithmeticOpDoc {
254+
let summary = "Scaled matrix multiply and accumulate";
255+
let description = [{
256+
Performs a scaled matrix multiply-accumulate using microscaling floating
257+
point (MXFP) formats. Each data operand has an associated per-group
258+
scale factor (E8M0 format). The computation is:
259+
260+
result = accumulator + (lhs_scale * lhs) @ (rhs_scale * rhs)
261+
262+
where scales are applied per group of K elements (group size is fixed
263+
at 32 by hardware).
264+
}] # baseDescription;
265+
266+
let arguments = !con((ins
267+
Arg<WaveTensorInRegister, "Left-hand side of the multiplication">:$lhs,
268+
Arg<WaveTensorInRegister, "Scale factors for the left-hand side">:$lhs_scale,
269+
Arg<WaveTensorInRegister, "Right-hand side of the multiplication">:$rhs,
270+
Arg<WaveTensorInRegister, "Scale factors for the right-hand side">:$rhs_scale,
271+
Arg<WaveTensorInRegister, "Accumulator for addition">:$accumulator,
272+
Arg<OptionalAttr<WaveMmaKindAttr>, "Kind of the scaled MMA intrinsic to target">:$kind
273+
), commonArguments);
274+
275+
let results = (outs
276+
Arg<WaveTensorInRegister, "Result">:$result
277+
);
278+
279+
let assemblyFormat =
280+
"$lhs `,` $lhs_scale `,` $rhs `,` $rhs_scale `,` $accumulator "
281+
# commonArgumentsSyntax # "attr-dict `:` functional-type(operands, results)";
282+
let hasVerifier = 1;
283+
284+
let extraClassDeclaration = [{
285+
/// Compute the expected elements per thread for a specific operand of this
286+
/// scaled MMA operation. Returns failure if no hardware constraints are
287+
/// available.
288+
llvm::FailureOr<unsigned> computeElementsPerThreadForOperand(unsigned operandIndex);
289+
}];
290+
}
291+
247292
//-----------------------------------------------------------------------------
248293
// Control flow operations
249294
//-----------------------------------------------------------------------------

0 commit comments

Comments
 (0)