Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 25 additions & 5 deletions include/flydsl/Dialect/FlyROCDL/IR/MmaAtom.td
Original file line number Diff line number Diff line change
Expand Up @@ -85,17 +85,37 @@ def FlyROCDL_MmaOpGFX1250_WMMA : FlyROCDL_MmaOp<"MmaOpGFX1250_WMMA", "gfx1250.wm
// false = unsigned, no clamp.
"bool":$signA,
"bool":$signB,
"bool":$clamp
"bool":$clamp,
// Intrinsic attributes forwarded to the ROCDL WMMA op: modC (I16
// C-operand modifier) and reuseA/reuseB (I1 operand-reuse scheduler
// hints). Default 0/false; elided from the assembly when at defaults.
DefaultValuedParameter<"int32_t", "0">:$modC,
DefaultValuedParameter<"bool", "false">:$reuseA,
DefaultValuedParameter<"bool", "false">:$reuseB
);
let assemblyFormat = "`<` custom<MNKDimensionList>($m, $n, $k) `,` `(` $elemTyA `,` $elemTyB `)` `->` $elemTyAcc `,` `signA` `=` $signA `,` `signB` `=` $signB `,` `clamp` `=` $clamp `>`";
let assemblyFormat = [{
`<` custom<MNKDimensionList>($m, $n, $k) `,` `(` $elemTyA `,` $elemTyB `)` `->` $elemTyAcc
`,` `signA` `=` $signA `,` `signB` `=` $signB `,` `clamp` `=` $clamp
(`,` `modC` `=` $modC^)?
(`,` `reuseA` `=` $reuseA^)?
(`,` `reuseB` `=` $reuseB^)? `>`
}];

let builders = [
// Back-compat: default sign/clamp to false (unsigned, no clamp).
// Back-compat: default sign/clamp/modC/reuse to false/0.
TypeBuilderWithInferredContext<(ins "int32_t":$m, "int32_t":$n, "int32_t":$k, "Type":$elemTyA, "Type":$elemTyB, "Type":$elemTyAcc), [{
return $_get(elemTyA.getContext(), m, n, k, elemTyA, elemTyB, elemTyAcc, /*signA=*/false, /*signB=*/false, /*clamp=*/false);
return $_get(elemTyA.getContext(), m, n, k, elemTyA, elemTyB, elemTyAcc,
/*signA=*/false, /*signB=*/false, /*clamp=*/false,
/*modC=*/0, /*reuseA=*/false, /*reuseB=*/false);
}]>,
TypeBuilderWithInferredContext<(ins "int32_t":$m, "int32_t":$n, "int32_t":$k, "Type":$elemTyA, "Type":$elemTyB, "Type":$elemTyAcc, "bool":$signA, "bool":$signB, "bool":$clamp), [{
return $_get(elemTyA.getContext(), m, n, k, elemTyA, elemTyB, elemTyAcc, signA, signB, clamp);
return $_get(elemTyA.getContext(), m, n, k, elemTyA, elemTyB, elemTyAcc,
signA, signB, clamp,
/*modC=*/0, /*reuseA=*/false, /*reuseB=*/false);
}]>,
TypeBuilderWithInferredContext<(ins "int32_t":$m, "int32_t":$n, "int32_t":$k, "Type":$elemTyA, "Type":$elemTyB, "Type":$elemTyAcc, "bool":$signA, "bool":$signB, "bool":$clamp, "int32_t":$modC, "bool":$reuseA, "bool":$reuseB), [{
return $_get(elemTyA.getContext(), m, n, k, elemTyA, elemTyB, elemTyAcc,
signA, signB, clamp, modC, reuseA, reuseB);
}]>
];
let genVerifyDecl = 1;
Expand Down
16 changes: 10 additions & 6 deletions lib/Bindings/Python/FlyROCDLExtension.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -66,17 +66,21 @@ struct PyMmaOpGFX1250_WMMAType : PyConcreteType<PyMmaOpGFX1250_WMMAType> {
c.def_static(
"get",
[](int32_t m, int32_t n, int32_t k, PyType &elemTyA, PyType &elemTyB, PyType &elemTyAcc,
bool signA, bool signB, bool clamp, DefaultingPyMlirContext context) {
bool signA, bool signB, bool clamp, int32_t modC, bool reuseA, bool reuseB,
DefaultingPyMlirContext context) {
return PyMmaOpGFX1250_WMMAType(
context->getRef(),
wrap(MmaOpGFX1250_WMMAType::get(m, n, k, unwrap(elemTyA), unwrap(elemTyB),
unwrap(elemTyAcc), signA, signB, clamp)));
context->getRef(), wrap(MmaOpGFX1250_WMMAType::get(
m, n, k, unwrap(elemTyA), unwrap(elemTyB), unwrap(elemTyAcc),
signA, signB, clamp, modC, reuseA, reuseB)));
},
"m"_a, "n"_a, "k"_a, "elem_ty_a"_a, "elem_ty_b"_a, "elem_ty_acc"_a, "sign_a"_a = false,
"sign_b"_a = false, "clamp"_a = false, nb::kw_only(), "context"_a = nb::none(),
"sign_b"_a = false, "clamp"_a = false, "mod_c"_a = 0, "reuse_a"_a = false,
"reuse_b"_a = false, nb::kw_only(), "context"_a = nb::none(),
"Create a MmaOpGFX1250_WMMAType with m, n, k dimensions and element types. "
"sign_a / sign_b / clamp are integer-only (iu4 / iu8) controls (signed operands / "
"accumulator saturation); they must be false on the float paths.");
"accumulator saturation); they must be false on the float paths. "
"mod_c (I16 C-operand modifier, default 0), reuse_a / reuse_b (operand-reuse "
"scheduler hints, default false) are forwarded to the ROCDL intrinsic.");
}
};

Expand Down
22 changes: 13 additions & 9 deletions lib/Dialect/FlyROCDL/GFX1250/MmaAtom.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,8 @@ Attribute MmaOpGFX1250_WMMAType::getThrValLayoutC() const {

LogicalResult MmaOpGFX1250_WMMAType::verify(function_ref<InFlightDiagnostic()> emitError, int32_t m,
int32_t n, int32_t k, Type elemTyA, Type elemTyB,
Type elemTyAcc, bool signA, bool signB, bool clamp) {
Type elemTyAcc, bool signA, bool signB, bool clamp,
int32_t modC, bool reuseA, bool reuseB) {
if (m != 16 || n != 16)
return emitError() << "GFX1250 WMMA requires M=N=16, got " << m << "x" << n;

Expand Down Expand Up @@ -235,19 +236,19 @@ enum class WmmaVariant { ModsAllReuse, ModsC, ModsABClamp, ModsIUClamp };
template <typename WmmaOp, WmmaVariant Variant>
static FailureOr<Value> emitWmmaSSA(OpBuilder &builder, Location loc, VectorType accTy, Value a,
Value b, Value c, bool signA = false, bool signB = false,
bool clamp = false) {
bool clamp = false, int32_t modC = 0, bool reuseA = false,
bool reuseB = false) {
Value res;
if constexpr (Variant == WmmaVariant::ModsAllReuse) {
res = WmmaOp::create(builder, loc, accTy, a, b, ROCDL::WMMACModifier::none, c,
/*reuseA=*/false, /*reuseB=*/false)
res = WmmaOp::create(builder, loc, accTy, a, b, static_cast<ROCDL::WMMACModifier>(modC), c,
reuseA, reuseB)
.getResult();
} else if constexpr (Variant == WmmaVariant::ModsC) {
res = WmmaOp::create(builder, loc, accTy, a, b, ROCDL::WMMACModifier::none, c,
/*reuseA=*/false, /*reuseB=*/false)
res = WmmaOp::create(builder, loc, accTy, a, b, static_cast<ROCDL::WMMACModifier>(modC), c,
reuseA, reuseB)
.getResult();
} else if constexpr (Variant == WmmaVariant::ModsABClamp) {
res = WmmaOp::create(builder, loc, accTy, signA, a, signB, b, c,
/*reuseA=*/false, /*reuseB=*/false, clamp)
res = WmmaOp::create(builder, loc, accTy, signA, a, signB, b, c, reuseA, reuseB, clamp)
.getResult();
} else {
static_assert(Variant == WmmaVariant::ModsIUClamp);
Expand Down Expand Up @@ -291,11 +292,14 @@ FailureOr<Value> MmaOpGFX1250_WMMAType::emitAtomCallSSA(OpBuilder &builder, Loca
bool signA = getSignA();
bool signB = getSignB();
bool clamp = getClamp();
int32_t modC = getModC();
bool reuseA = getReuseA();
bool reuseB = getReuseB();

#define DISPATCH_WMMA_SSA(M_, K_, PRED, OP, VARIANT) \
if (m == M_ && n == M_ && k == K_ && (PRED)) { \
return emitWmmaSSA<ROCDL::OP, WmmaVariant::VARIANT>(builder, loc, accTy, a, b, c, signA, \
signB, clamp); \
signB, clamp, modC, reuseA, reuseB); \
}

#define DISPATCH_WMMA_SSA_FP8(K_, ACC_PRED, ACC_PREFIX) \
Expand Down
21 changes: 14 additions & 7 deletions python/flydsl/expr/rocdl/universal.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,15 +117,19 @@ def MFMA(m, n, k, elem_ty_ab, elem_ty_acc=None):
def WMMA(m, n, k, elem_ty_ab, elem_ty_acc=None, **kwargs):
"""Create an arch-appropriate WMMA atom.

Supported kwargs (integer paths only — iu8 / iu4):
sign_a (bool, default False): treat A operand as signed.
sign_b (bool, default False): treat B operand as signed.
clamp (bool, default False): saturate integer accumulator.
Supported kwargs:
sign_a (bool, default False): treat A operand as signed (iu8/iu4 only).
sign_b (bool, default False): treat B operand as signed (iu8/iu4 only).
clamp (bool, default False): saturate integer accumulator (iu8/iu4 only).
mod_c (int, default 0): I16 C-operand modifier (gfx1250 only).
reuse_a (bool, default False): operand-reuse scheduler hint (gfx1250 only).
reuse_b (bool, default False): operand-reuse scheduler hint (gfx1250 only).
Forwarded to the arch-specific WMMA atom (MmaOpGFX11_WMMAType on gfx11,
MmaOpGFX120X_WMMAType on gfx120x, MmaOpGFX1250_WMMAType on gfx1250); the
atom's verify() rejects them on the float (fp16/bf16/fp8) paths, where the
intrinsic has no such operands. Future WMMA ops for new architectures
should extend kwargs here rather than growing the positional signature.
atom's verify() rejects sign_a/sign_b/clamp on the float (fp16/bf16/fp8)
paths, where the intrinsic has no such operands. Future WMMA ops for new
architectures should extend kwargs here rather than growing the positional
signature.
"""
ty_ab = elem_ty_ab.ir_type if hasattr(elem_ty_ab, "ir_type") else elem_ty_ab
if elem_ty_acc is None:
Expand Down Expand Up @@ -155,6 +159,9 @@ def WMMA(m, n, k, elem_ty_ab, elem_ty_acc=None, **kwargs):
sign_a=bool(kwargs.get("sign_a", False)),
sign_b=bool(kwargs.get("sign_b", False)),
clamp=bool(kwargs.get("clamp", False)),
mod_c=int(kwargs.get("mod_c", 0)),
reuse_a=bool(kwargs.get("reuse_a", False)),
reuse_b=bool(kwargs.get("reuse_b", False)),
)
if arch.startswith("gfx120"):
return MmaOpGFX120X_WMMAType.get(
Expand Down
23 changes: 23 additions & 0 deletions tests/mlir/Conversion/wmma_gfx1250.mlir
Original file line number Diff line number Diff line change
Expand Up @@ -49,3 +49,26 @@ func.func @test_wmma_iu4_signed_clamp(
fly.mma_atom_call(%atom, %d, %a, %b, %c) : (!fly.mma_atom<!fly_rocdl.gfx1250.wmma<16x16x32, (i4, i4) -> i32, signA = true, signB = true, clamp = true>>, !fly.memref<i32, register, 8:1>, !fly.memref<i4, register, 16:1>, !fly.memref<i4, register, 16:1>, !fly.memref<i32, register, 8:1>) -> ()
return
}

// -----

// bf16 WMMA with modC and reuse controls: verifies that modC / reuseA /
// reuseB are forwarded to the emitted rocdl.wmma op on the bf16 path.

// CHECK-LABEL: @test_wmma_bf16_modc_reuse
func.func @test_wmma_bf16_modc_reuse(
%atom: !fly.mma_atom<!fly_rocdl.gfx1250.wmma<16x16x32, (bf16, bf16) -> f32, signA = false, signB = false, clamp = false, modC = 1, reuseA = true, reuseB = true>>) {
%lay_ab = fly.static : !fly.layout<16:1>
%lay_cd = fly.static : !fly.layout<8:1>
%d = fly.memref.alloca(%lay_cd) : (!fly.layout<8:1>) -> !fly.memref<f32, register, 8:1>
%a = fly.memref.alloca(%lay_ab) : (!fly.layout<16:1>) -> !fly.memref<bf16, register, 16:1>
%b = fly.memref.alloca(%lay_ab) : (!fly.layout<16:1>) -> !fly.memref<bf16, register, 16:1>
%c = fly.memref.alloca(%lay_cd) : (!fly.layout<8:1>) -> !fly.memref<f32, register, 8:1>

// CHECK: rocdl.wmma.f32.16x16x32.bf16
// CHECK-SAME: modC = neg
// CHECK-SAME: reuseA = true
// CHECK-SAME: reuseB = true
fly.mma_atom_call(%atom, %d, %a, %b, %c) : (!fly.mma_atom<!fly_rocdl.gfx1250.wmma<16x16x32, (bf16, bf16) -> f32, signA = false, signB = false, clamp = false, modC = 1, reuseA = true, reuseB = true>>, !fly.memref<f32, register, 8:1>, !fly.memref<bf16, register, 16:1>, !fly.memref<bf16, register, 16:1>, !fly.memref<f32, register, 8:1>) -> ()
return
}
20 changes: 20 additions & 0 deletions tests/unit/test_gfx1250_atoms.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,26 @@ def test_wmma_scale_type_roundtrip():
assert ir.Type.parse(str(t_reuse)) == t_reuse


def test_wmma_type_modc_reuse_roundtrip():
with _ctx(), ir.Location.unknown():
from flydsl._mlir._mlir_libs._mlirDialectsFlyROCDL import MmaOpGFX1250_WMMAType
from flydsl._mlir.dialects import fly_rocdl # noqa: F401

bf16 = ir.BF16Type.get()
f32 = ir.F32Type.get()

t_default = MmaOpGFX1250_WMMAType.get(16, 16, 32, bf16, bf16, f32)
assert "gfx1250.wmma<" in str(t_default)
# Defaults (modC=0, reuseA/reuseB=false) are elided from the printed form.
assert "modC" not in str(t_default)
assert ir.Type.parse(str(t_default)) == t_default

t_modc = MmaOpGFX1250_WMMAType.get(16, 16, 32, bf16, bf16, f32, mod_c=1, reuse_a=True, reuse_b=True)
assert "modC = 1, reuseA = true, reuseB = true" in str(t_modc)
assert ir.Type.parse(str(t_modc)) == t_modc
assert t_modc != t_default


def test_tdm2d_type_roundtrip():
with _ctx(), ir.Location.unknown():
from flydsl._mlir.dialects import fly_rocdl # noqa: F401
Expand Down
Loading
Loading