diff --git a/NOTICE b/NOTICE new file mode 100644 index 000000000..daeda5837 --- /dev/null +++ b/NOTICE @@ -0,0 +1,41 @@ +FlyDSL +Copyright (c) 2025 FlyDSL Project Contributors + +This product includes software developed as part of the FlyDSL project, +licensed under the Apache License, Version 2.0 (see LICENSE). + +------------------------------------------------------------------------ +Third-party components +------------------------------------------------------------------------ + +Triton (https://github.com/triton-lang/triton) — MIT License + +The gfx1250 TDM (Tensor Data Mover) copy-atom lowering in + lib/Dialect/FlyROCDL/GFX1250/CopyAtom.cpp +is a port of the AMD gfx1250 TDM descriptor bitfield layout, the N-D warp +distribution, and the row-gather descriptor packing from Triton's AMD backend +(third_party/amd/lib/TritonAMDGPUToLLVM/TDMUtility.cpp and +.../backend/include/TDMCommon.h: createTDMDescriptor / fillTDMDescriptor / +tdmGetWarpDistribution). Triton is used under the MIT License: + + Copyright 2018-2020 Philippe Tillet + Copyright 2020-2022 OpenAI + + Permission is hereby granted, free of charge, to any person obtaining + a copy of this software and associated documentation files (the + "Software"), to deal in the Software without restriction, including + without limitation the rights to use, copy, modify, merge, publish, + distribute, sublicense, and/or sell copies of the Software, and to + permit persons to whom the Software is furnished to do so, subject to + the following conditions: + + The above copyright notice and this permission notice shall be + included in all copies or substantial portions of the Software. + + THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, + EXPRESS OR IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF + MERCHANTABILITY, FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. + IN NO EVENT SHALL THE AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY + CLAIM, DAMAGES OR OTHER LIABILITY, WHETHER IN AN ACTION OF CONTRACT, + TORT OR OTHERWISE, ARISING FROM, OUT OF OR IN CONNECTION WITH THE + SOFTWARE OR THE USE OR OTHER DEALINGS IN THE SOFTWARE. diff --git a/README.md b/README.md index e2ef4d220..0a96273ea 100644 --- a/README.md +++ b/README.md @@ -406,7 +406,7 @@ FlyDSL's design is inspired by ideas from several projects: - [NVIDIA CUTLASS](https://github.com/NVIDIA/cutlass) — CuTe layout algebra concepts (BSD-3-Clause parts only; no EULA-licensed code was referenced) - [ROCm Composable Kernel](https://github.com/ROCm/composable_kernel) — tile-based kernel design patterns for AMD GPUs - [ROCm AIter](https://github.com/ROCm/aiter) — test infrastructure and performance comparison baselines (MIT) -- [Triton](https://github.com/triton-lang/triton) — Python DSL for GPU kernel authoring +- [Triton](https://github.com/triton-lang/triton) — Python DSL for GPU kernel authoring; the gfx1250 TDM copy-atom descriptor packing is ported from Triton's AMD backend (MIT, see [`NOTICE`](NOTICE)) - [HipKittens](https://github.com/HazyResearch/HipKittens) — minimal, opinionated C++ embedded primitives for fast AMD AI kernels (part of the ThunderKittens family) ## 📄 License diff --git a/examples/05-gather_scatter.py b/examples/05-gather_scatter.py index 235e4275a..de04eaf1a 100644 --- a/examples/05-gather_scatter.py +++ b/examples/05-gather_scatter.py @@ -21,6 +21,18 @@ offset tensor : (TV, Rest...) pred tensor : (TV, Rest...) optional +The same ``fx.gather`` / ``fx.scatter`` entry points also drive hardware +whole-tile gather on gfx1250: pass a TDM gather atom +(``fx.rocdl.make_tdm_gather_atom(src)``, ``copy_rank == 2``) instead of a +per-element ``UniversalCopy`` atom and the call lowers to a single +``rocdl.tensor.load.to.lds`` / ``store.from.lds`` gather. In that mode the +``offset`` argument is the row-index operand (an i16/i32 tensor whose element +type sets the index width and whose length sets the row count), ``base_iter`` +only supplies the global shape/direction token (the base pointer is atom +state), and no per-instance loop is emitted. See +``tests/mlir/Conversion/tdm_gather_gfx1250.mlir`` for the lowering and +``fx.rocdl.make_tdm_gather_atom`` for the atom builder. The example kernels +below use the arch-neutral ``UniversalCopy`` path so they run everywhere. """ import torch diff --git a/include/flydsl/Dialect/Fly/IR/FlyInterfaces.td b/include/flydsl/Dialect/Fly/IR/FlyInterfaces.td index 9c618053c..965719900 100644 --- a/include/flydsl/Dialect/Fly/IR/FlyInterfaces.td +++ b/include/flydsl/Dialect/Fly/IR/FlyInterfaces.td @@ -80,7 +80,9 @@ def Fly_CopyOpTypeInterface : TypeInterface<"CopyOpTypeInterface"> { InterfaceMethod<"", "::mlir::Attribute", "getThrBitLayoutDst", (ins)>, InterfaceMethod<"", "::mlir::Attribute", "getThrBitLayoutRef", (ins)>, InterfaceMethod< - "Emit the lowering IR for a copy_atom_call with this CopyOp type.", + "Emit the lowering IR for a copy_atom_call with this CopyOp type. " + "`indicesMemTy`/`indices` carry the optional gather/scatter row-index " + "operand (null when absent); atoms that do not gather ignore them.", "::mlir::LogicalResult", "emitAtomCall", (ins "::mlir::OpBuilder &":$builder, @@ -88,11 +90,15 @@ def Fly_CopyOpTypeInterface : TypeInterface<"CopyOpTypeInterface"> { "::mlir::Type":$copyAtomTy, "::mlir::Type":$srcMemTy, "::mlir::Type":$dstMemTy, + "::mlir::Type":$indicesMemTy, "::mlir::Value":$atomVal, "::mlir::Value":$src, - "::mlir::Value":$dst)>, + "::mlir::Value":$dst, + "::mlir::Value":$indices)>, InterfaceMethod< - "Emit the lowering IR for a predicated copy_atom_call with this CopyOp type.", + "Emit the lowering IR for a predicated copy_atom_call with this CopyOp type. " + "`indicesMemTy`/`indices` carry the optional gather/scatter row-index " + "operand (null when absent); atoms that do not gather ignore them.", "::mlir::LogicalResult", "emitAtomCall", (ins "::mlir::OpBuilder &":$builder, @@ -100,10 +106,12 @@ def Fly_CopyOpTypeInterface : TypeInterface<"CopyOpTypeInterface"> { "::mlir::Type":$copyAtomTy, "::mlir::Type":$srcMemTy, "::mlir::Type":$dstMemTy, + "::mlir::Type":$indicesMemTy, "::mlir::Type":$predMemTy, "::mlir::Value":$atomVal, "::mlir::Value":$src, "::mlir::Value":$dst, + "::mlir::Value":$indices, "::mlir::Value":$pred)>, InterfaceMethod< "Emit SSA-form copy: reads src, returns vector result.", diff --git a/include/flydsl/Dialect/Fly/IR/FlyOps.td b/include/flydsl/Dialect/Fly/IR/FlyOps.td index 20a5f5b43..33e5d8380 100644 --- a/include/flydsl/Dialect/Fly/IR/FlyOps.td +++ b/include/flydsl/Dialect/Fly/IR/FlyOps.td @@ -357,8 +357,13 @@ def Fly_AtomSetValueOp : Fly_Op<"atom.set_value", [Pure, DeclareOpInterfaceMetho let assemblyFormat = "`(` $atom `,` $field `,` $value `)` attr-dict `:` functional-type(operands, results)"; } -def Fly_CopyAtomCall : Fly_Op<"copy_atom_call"> { - let arguments = (ins Fly_CopyAtom:$copyAtom, Fly_MemRef:$src, Fly_MemRef:$dst, Optional:$pred); +def Fly_CopyAtomCall : Fly_Op<"copy_atom_call", [AttrSizedOperandSegments]> { + let arguments = (ins Fly_CopyAtom:$copyAtom, Fly_MemRef:$src, Fly_MemRef:$dst, + Optional:$indices, Optional:$pred); + // `indices` is keyword-led (no leading comma) so it stays distinguishable from + // the bare comma-led optional `pred`; both may appear, indices then pred. + let assemblyFormat = "`(` $copyAtom `,` $src `,` $dst (`indices` `=` $indices^)? (`,` $pred^)? `)` " + "attr-dict `:` functional-type(operands, results)"; } def Fly_MmaAtomCall : Fly_Op<"mma_atom_call"> { let arguments = (ins Fly_MmaAtom:$mmaAtom, Fly_MemRef:$d, Fly_MemRef:$a, Fly_MemRef:$b, Fly_MemRef:$c); @@ -414,9 +419,13 @@ def Fly_MmaMakeFragmentOp : Fly_Op<"mma.make_fragment", [Pure, DeclareOpInterfac let assemblyFormat = "`(` $operand_id `,` $tiled_mma `,` $input (`,` `stages` `=` $stages^)? `)` attr-dict `:` functional-type(operands, results)"; } -def Fly_CopyOp : Fly_Op<"copy"> { - let arguments = (ins AnyType:$copyAtom, Fly_MemRef:$src, Fly_MemRef:$dst, Optional:$pred); - let assemblyFormat = "`(` $copyAtom `,` $src `,` $dst (`,` $pred^)? `)` attr-dict `:` functional-type(operands, results)"; +def Fly_CopyOp : Fly_Op<"copy", [AttrSizedOperandSegments]> { + let arguments = (ins AnyType:$copyAtom, Fly_MemRef:$src, Fly_MemRef:$dst, + Optional:$indices, Optional:$pred); + // `indices` is keyword-led (no leading comma) so it stays distinguishable from + // the bare comma-led optional `pred`; both may appear, indices then pred. + let assemblyFormat = "`(` $copyAtom `,` $src `,` $dst (`indices` `=` $indices^)? (`,` $pred^)? `)` " + "attr-dict `:` functional-type(operands, results)"; } def Fly_GemmOp : Fly_Op<"gemm"> { let arguments = (ins AnyType:$mmaAtom, Fly_MemRef:$d, Fly_MemRef:$a, Fly_MemRef:$b, Fly_MemRef:$c, diff --git a/include/flydsl/Dialect/Fly/IR/FlyTypeDefs.td b/include/flydsl/Dialect/Fly/IR/FlyTypeDefs.td index 4ce9cc932..95f2c1d66 100644 --- a/include/flydsl/Dialect/Fly/IR/FlyTypeDefs.td +++ b/include/flydsl/Dialect/Fly/IR/FlyTypeDefs.td @@ -269,10 +269,11 @@ def Fly_CopyAtom : Fly_Type<"CopyAtom", "copy_atom", [ ::mlir::Attribute getThrValLayoutRef(); ::mlir::LogicalResult emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTy, - Type srcMemTy, Type dstMemTy, Value atomVal, Value src, Value dst) const; + Type srcMemTy, Type dstMemTy, Type indicesMemTy, Value atomVal, + Value src, Value dst, Value indices) const; ::mlir::LogicalResult emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTy, - Type srcMemTy, Type dstMemTy, Type predMemTy, Value atomVal, Value src, Value dst, - Value pred) const; + Type srcMemTy, Type dstMemTy, Type indicesMemTy, Type predMemTy, + Value atomVal, Value src, Value dst, Value indices, Value pred) const; ::mlir::FailureOr<::mlir::Value> emitAtomCallSSA(OpBuilder &builder, Location loc, Type resultTy, Type copyAtomTy, Type srcTy, Type dstTy, diff --git a/include/flydsl/Dialect/FlyROCDL/IR/Atom.td b/include/flydsl/Dialect/FlyROCDL/IR/Atom.td index 9d1f44fe5..d6cabc4e2 100644 --- a/include/flydsl/Dialect/FlyROCDL/IR/Atom.td +++ b/include/flydsl/Dialect/FlyROCDL/IR/Atom.td @@ -31,6 +31,14 @@ def FlyROCDL_AtomStateField : I32EnumAttr<"AtomStateField", "", [ let cppNamespace = FlyROCDL_Dialect.cppNamespace; } +// gfx1250 TDM copy mode: whole-tile contiguous move (tiled) vs row gather/scatter +// selected by an explicit row-index operand (gather). Direction (load vs store) is +// still inferred from the operand address spaces. +def FlyROCDL_TdmMode : FlyROCDL_I32EnumAttr<"TdmMode", "gfx1250 TDM copy mode", [ + I32EnumAttrCase<"Tiled", 0, "tiled">, + I32EnumAttrCase<"Gather", 1, "gather"> +]>; + include "flydsl/Dialect/FlyROCDL/IR/CopyAtom.td" include "flydsl/Dialect/FlyROCDL/IR/MmaAtom.td" diff --git a/include/flydsl/Dialect/FlyROCDL/IR/CopyAtom.td b/include/flydsl/Dialect/FlyROCDL/IR/CopyAtom.td index 0f5961f9e..bb4b32d90 100644 --- a/include/flydsl/Dialect/FlyROCDL/IR/CopyAtom.td +++ b/include/flydsl/Dialect/FlyROCDL/IR/CopyAtom.td @@ -62,6 +62,11 @@ def FlyROCDL_CopyOpLdsReadTranspose : FlyROCDL_CopyOp<"CopyOpCDNA4LdsReadTranspo // dims 0..rank-2 (innermost stride is assumed 1). Unset falls back to the // tile memref's static layout stride. // `imm_offset` (default 0): i64 byte offset added to base (K-loop tile bump). +// +// `mode` selects the descriptor packing: `tiled` moves a contiguous N-D tile; +// `gather` (rank 2 only) reads a set of rows selected by an `indices` memref +// operand on the copy_atom_call, with the row-index width taken from that operand's +// element type (i16 or i32) and the row count from its static extent. //===----------------------------------------------------------------------===// // TDM moves the whole N-D tile in one DMA, so it overrides `getCopyRank` (== rank) @@ -79,12 +84,14 @@ def FlyROCDL_CopyOpGFX1250TDM : FlyROCDL_StatefulCopyOp<"CopyOpGFX1250TDM", "gfx // Descriptor config bits (GROUP1 sgpr0): atomic_barrier_enable [18] sets the // HW auto-barrier bit; early_timeout [21] is a multicast-load GL1 knob. "bool":$atomicBarrier, - "bool":$earlyTimeout + "bool":$earlyTimeout, + // Copy mode: tiled (contiguous N-D tile) or gather (row gather/scatter). + EnumParameter:$mode ); let assemblyFormat = [{ `<` `rank` `=` $rank `,` `warps` `=` $numWarps `,` `pad` `=` $padInterval `,` $padAmount `,` `cache` `=` $cacheModifier `,` `barrier` `=` $atomicBarrier - `,` `timeout` `=` $earlyTimeout `>` + `,` `timeout` `=` $earlyTimeout `,` `mode` `=` $mode `>` }]; let genVerifyDecl = 1; } diff --git a/lib/Bindings/Python/FlyExtension.cpp b/lib/Bindings/Python/FlyExtension.cpp index f974dae9c..1bba9e887 100644 --- a/lib/Bindings/Python/FlyExtension.cpp +++ b/lib/Bindings/Python/FlyExtension.cpp @@ -777,6 +777,11 @@ struct PyCopyAtomType : PyConcreteType { return wrap(self.toCppType().getCopyOp()); }); c.def_prop_ro("val_bits", [](PyCopyAtomType &self) { return self.toCppType().getValBits(); }); + // Number of memref dims the atom transfers per copy_atom_call (whole-tile atoms + // like TDM return N); the DSL gather/scatter routing keys on this. + c.def_prop_ro("copy_rank", [](PyCopyAtomType &self) -> unsigned { + return cast(self.toCppType().getCopyOp()).getCopyRank(); + }); c.def_prop_ro("thr_layout", [](PyCopyAtomType &self) -> MlirType { return wrap(LayoutType::get(cast(self.toCppType().getThrLayout()))); }); diff --git a/lib/Bindings/Python/FlyROCDLExtension.cpp b/lib/Bindings/Python/FlyROCDLExtension.cpp index c7c63895a..da3aeffa9 100644 --- a/lib/Bindings/Python/FlyROCDLExtension.cpp +++ b/lib/Bindings/Python/FlyROCDLExtension.cpp @@ -186,20 +186,22 @@ struct PyCopyOpGFX1250TDMType : PyConcreteType { c.def_static( "get", [](int32_t rank, int32_t numWarps, int32_t padInterval, int32_t padAmount, - int32_t cacheModifier, bool atomicBarrier, bool earlyTimeout, + int32_t cacheModifier, bool atomicBarrier, bool earlyTimeout, int32_t mode, DefaultingPyMlirContext context) { MLIRContext *ctx = unwrap(context.get()->get()); return PyCopyOpGFX1250TDMType( context->getRef(), wrap(CopyOpGFX1250TDMType::get(ctx, rank, numWarps, padInterval, padAmount, - cacheModifier, atomicBarrier, earlyTimeout))); + cacheModifier, atomicBarrier, earlyTimeout, + static_cast(mode)))); }, "rank"_a, "num_warps"_a, "pad_interval"_a = 0, "pad_amount"_a = 0, "cache_modifier"_a = 0, - "atomic_barrier"_a = false, "early_timeout"_a = false, nb::kw_only(), + "atomic_barrier"_a = false, "early_timeout"_a = false, "mode"_a = 0, nb::kw_only(), "context"_a = nb::none(), "Create a CopyOpGFX1250TDMType (N-D TDM Global<->LDS copy) with tensor rank (1-5), " - "warp count, optional LDS padding (interval/amount in elements), cache modifier, and " - "the descriptor atomic_barrier / early_timeout config bits (default false)"); + "warp count, optional LDS padding (interval/amount in elements), cache modifier, " + "the descriptor atomic_barrier / early_timeout config bits (default false), and " + "mode (0 = tiled, 1 = gather)"); } }; diff --git a/lib/Conversion/FlyToROCDL/FlyToROCDL.cpp b/lib/Conversion/FlyToROCDL/FlyToROCDL.cpp index 0200b5b0a..c4a2c4e20 100644 --- a/lib/Conversion/FlyToROCDL/FlyToROCDL.cpp +++ b/lib/Conversion/FlyToROCDL/FlyToROCDL.cpp @@ -534,6 +534,7 @@ class CopyAtomCallLowering : public OpConversionPattern { Value copyAtomVal = adaptor.getCopyAtom(); Value src = adaptor.getSrc(); Value dst = adaptor.getDst(); + Value indices = adaptor.getIndices(); Value pred = adaptor.getPred(); auto srcMemTy = dyn_cast(op.getSrc().getType()); @@ -546,6 +547,14 @@ class CopyAtomCallLowering : public OpConversionPattern { Location loc = op.getLoc(); + // Optional gather/scatter row-index operand; null when absent. + Type indicesMemTy = nullptr; + if (indices) { + indicesMemTy = dyn_cast(op.getIndices().getType()); + if (!indicesMemTy) + return rewriter.notifyMatchFailure(op, "indices is not a MemRef type"); + } + Type predMemTy = nullptr; if (pred) { predMemTy = dyn_cast(op.getPred().getType()); @@ -554,12 +563,12 @@ class CopyAtomCallLowering : public OpConversionPattern { } if (pred) { - if (failed(copyAtom.emitAtomCall(rewriter, loc, copyAtomType, srcMemTy, dstMemTy, predMemTy, - copyAtomVal, src, dst, pred))) + if (failed(copyAtom.emitAtomCall(rewriter, loc, copyAtomType, srcMemTy, dstMemTy, indicesMemTy, + predMemTy, copyAtomVal, src, dst, indices, pred))) return failure(); } else { - if (failed(copyAtom.emitAtomCall(rewriter, loc, copyAtomType, srcMemTy, dstMemTy, copyAtomVal, - src, dst))) + if (failed(copyAtom.emitAtomCall(rewriter, loc, copyAtomType, srcMemTy, dstMemTy, indicesMemTy, + copyAtomVal, src, dst, indices))) return failure(); } rewriter.eraseOp(op); diff --git a/lib/Dialect/Fly/IR/FlyTypeDefs.cpp b/lib/Dialect/Fly/IR/FlyTypeDefs.cpp index efeada8ee..7cb34ad5e 100644 --- a/lib/Dialect/Fly/IR/FlyTypeDefs.cpp +++ b/lib/Dialect/Fly/IR/FlyTypeDefs.cpp @@ -450,18 +450,20 @@ Attribute CopyAtomType::getThrValLayoutRef() { } LogicalResult CopyAtomType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTy, - Type srcMemTy, Type dstMemTy, Value atomVal, Value src, - Value dst) const { + Type srcMemTy, Type dstMemTy, Type indicesMemTy, + Value atomVal, Value src, Value dst, Value indices) const { return cast(getCopyOp()) - .emitAtomCall(builder, loc, copyAtomTy, srcMemTy, dstMemTy, atomVal, src, dst); + .emitAtomCall(builder, loc, copyAtomTy, srcMemTy, dstMemTy, indicesMemTy, atomVal, src, dst, + indices); } LogicalResult CopyAtomType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTy, - Type srcMemTy, Type dstMemTy, Type predMemTy, - Value atomVal, Value src, Value dst, Value pred) const { + Type srcMemTy, Type dstMemTy, Type indicesMemTy, + Type predMemTy, Value atomVal, Value src, Value dst, + Value indices, Value pred) const { return cast(getCopyOp()) - .emitAtomCall(builder, loc, copyAtomTy, srcMemTy, dstMemTy, predMemTy, atomVal, src, dst, - pred); + .emitAtomCall(builder, loc, copyAtomTy, srcMemTy, dstMemTy, indicesMemTy, predMemTy, atomVal, + src, dst, indices, pred); } FailureOr CopyAtomType::emitAtomCallSSA(OpBuilder &builder, Location loc, Type resultTy, diff --git a/lib/Dialect/Fly/IR/FlyUniversalOps.cpp b/lib/Dialect/Fly/IR/FlyUniversalOps.cpp index 8f470168d..faf9dd1fa 100644 --- a/lib/Dialect/Fly/IR/FlyUniversalOps.cpp +++ b/lib/Dialect/Fly/IR/FlyUniversalOps.cpp @@ -180,8 +180,9 @@ FailureOr CopyOpUniversalCopyType::emitAtomCallSSA(OpBuilder &builder, Lo LogicalResult CopyOpUniversalCopyType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, Value src, - Value dst) const { + Type dstMemTyArg, Type /*indicesMemTyArg*/, + Value atomVal, Value src, Value dst, + Value /*indices*/) const { auto srcMemTy = cast(srcMemTyArg); auto dstMemTy = cast(dstMemTyArg); @@ -201,15 +202,16 @@ LogicalResult CopyOpUniversalCopyType::emitAtomCall(OpBuilder &builder, Location LogicalResult CopyOpUniversalCopyType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, - Value atomVal, Value src, Value dst, - Value pred) const { + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, Value src, + Value dst, Value indices, Value pred) const { auto predMemTy = cast(predMemTyArg); Value predVal = LLVM::LoadOp::create(builder, loc, predMemTy.getElemTy(), pred); auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } static std::optional convertAtomicOp(AtomicOp binOp, bool isFloat) { @@ -283,8 +285,9 @@ FailureOr CopyOpUniversalAtomicType::emitAtomCallSSA( LogicalResult CopyOpUniversalAtomicType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, Value src, - Value dst) const { + Type dstMemTyArg, Type /*indicesMemTyArg*/, + Value atomVal, Value src, Value dst, + Value /*indices*/) const { auto srcMemTy = cast(srcMemTyArg); auto srcSSATy = fly::RegMem2SSAType(srcMemTy, /*llvmCompatibleType=*/true); Value srcVal = LLVM::LoadOp::create(builder, loc, srcSSATy, src); @@ -297,15 +300,16 @@ LogicalResult CopyOpUniversalAtomicType::emitAtomCall(OpBuilder &builder, Locati LogicalResult CopyOpUniversalAtomicType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, - Value atomVal, Value src, Value dst, - Value pred) const { + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, Value src, + Value dst, Value indices, Value pred) const { auto predMemTy = cast(predMemTyArg); Value predVal = LLVM::LoadOp::create(builder, loc, predMemTy.getElemTy(), pred); auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } FailureOr MmaOpUniversalFMAType::emitAtomCallSSA(OpBuilder &builder, Location loc, diff --git a/lib/Dialect/Fly/Transforms/LayoutLowering.cpp b/lib/Dialect/Fly/Transforms/LayoutLowering.cpp index b58a710d5..bee5c743c 100644 --- a/lib/Dialect/Fly/Transforms/LayoutLowering.cpp +++ b/lib/Dialect/Fly/Transforms/LayoutLowering.cpp @@ -2168,6 +2168,7 @@ class ExpandCopyOpLowering : public OpRewritePattern { Value src = op.getSrc(); Value dst = op.getDst(); + Value indices = op.getIndices(); Value pred = op.getPred(); auto srcMemRefTy = cast(src.getType()); @@ -2209,13 +2210,18 @@ class ExpandCopyOpLowering : public OpRewritePattern { if (srcRank != static_cast(copyRank)) return rewriter.notifyMatchFailure( op, "whole-tile copy atom requires a memref of its exact copy rank"); - CopyAtomCall::create(rewriter, loc, copyAtomVal, src, dst, pred); + CopyAtomCall::create(rewriter, loc, copyAtomVal, src, dst, indices, pred); rewriter.eraseOp(op); return success(); } } } + // The gather/scatter row-index operand is only meaningful for whole-tile + // atoms handled above; the per-element decomposition below cannot carry it. + if (indices) + return rewriter.notifyMatchFailure(op, "indices operand requires a whole-tile copy atom"); + if (pred && predLayoutAttr.rank() == srcRank - 1) { LayoutBuilder builder(rewriter, loc); LayoutAttr unitAttr = LayoutAttr::get(ctx, IntTupleAttr::getLeafStatic(ctx, 1), @@ -2234,7 +2240,8 @@ class ExpandCopyOpLowering : public OpRewritePattern { if (srcLayoutAttr.getShape().isLeaf()) { Value srcDecomposition = DecompositionOp::create(rewriter, loc, src); Value dstDecomposition = DecompositionOp::create(rewriter, loc, dst); - CopyAtomCall::create(rewriter, loc, copyAtomVal, srcDecomposition, dstDecomposition, pred); + CopyAtomCall::create(rewriter, loc, copyAtomVal, srcDecomposition, dstDecomposition, + /*indices=*/Value{}, pred); rewriter.eraseOp(op); return success(); } @@ -2242,7 +2249,8 @@ class ExpandCopyOpLowering : public OpRewritePattern { Value dstUnwrapped = GetOp::create(rewriter, loc, dst, ArrayRef{0}); Value predUnwrapped = pred ? GetOp::create(rewriter, loc, pred, ArrayRef{0}) : nullptr; - CopyOp::create(rewriter, loc, copyAtomVal, srcUnwrapped, dstUnwrapped, predUnwrapped); + CopyOp::create(rewriter, loc, copyAtomVal, srcUnwrapped, dstUnwrapped, /*indices=*/Value{}, + predUnwrapped); rewriter.eraseOp(op); return success(); } @@ -2276,7 +2284,7 @@ class ExpandCopyOpLowering : public OpRewritePattern { if (pred) predSlice = SliceOp::create(rewriter, loc, predGrouped, coord); - CopyOp::create(rewriter, loc, copyAtomVal, srcSlice, dstSlice, predSlice); + CopyOp::create(rewriter, loc, copyAtomVal, srcSlice, dstSlice, /*indices=*/Value{}, predSlice); } rewriter.eraseOp(op); return success(); diff --git a/lib/Dialect/FlyROCDL/CDNA3/CopyAtom.cpp b/lib/Dialect/FlyROCDL/CDNA3/CopyAtom.cpp index 4b16116cb..4225f0f41 100644 --- a/lib/Dialect/FlyROCDL/CDNA3/CopyAtom.cpp +++ b/lib/Dialect/FlyROCDL/CDNA3/CopyAtom.cpp @@ -157,8 +157,9 @@ FailureOr CopyOpCDNA3BufferCopyType::emitAtomCallSSA( LogicalResult CopyOpCDNA3BufferCopyType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, Value src, - Value dst) const { + Type dstMemTyArg, Type /*indicesMemTyArg*/, + Value atomVal, Value src, Value dst, + Value /*indices*/) const { auto srcMemTy = cast(srcMemTyArg); auto dstMemTy = cast(dstMemTyArg); @@ -188,16 +189,17 @@ LogicalResult CopyOpCDNA3BufferCopyType::emitAtomCall(OpBuilder &builder, Locati LogicalResult CopyOpCDNA3BufferCopyType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, - Value atomVal, Value src, Value dst, - Value pred) const { + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, Value src, + Value dst, Value indices, Value pred) const { OpBuilder::InsertionGuard guard(builder); auto predMemTy = cast(predMemTyArg); Value predVal = LLVM::LoadOp::create(builder, loc, predMemTy.getElemTy(), pred); auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } // --- CopyOpCDNA3BufferCopyLDS --- @@ -267,7 +269,8 @@ FailureOr CopyOpCDNA3BufferCopyLDSType::emitAtomCallSSA(OpBuilder &builde Type srcTyArg, Type dstTyArg, Value atomVal, Value src, Value dst) const { - if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, atomVal, src, dst))) + if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, /*indicesMemTy=*/Type{}, + atomVal, src, dst, /*indices=*/Value{}))) return failure(); return Value{}; } @@ -275,16 +278,17 @@ FailureOr CopyOpCDNA3BufferCopyLDSType::emitAtomCallSSA(OpBuilder &builde FailureOr CopyOpCDNA3BufferCopyLDSType::emitAtomCallSSA( OpBuilder &builder, Location loc, Type resultTy, Type copyAtomTyArg, Type srcTyArg, Type dstTyArg, Type predTyArg, Value atomVal, Value src, Value dst, Value pred) const { - if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, predTyArg, atomVal, src, - dst, pred))) + if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, /*indicesMemTy=*/Type{}, + predTyArg, atomVal, src, dst, /*indices=*/Value{}, pred))) return failure(); return Value{}; } LogicalResult CopyOpCDNA3BufferCopyLDSType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, Value src, - Value dst) const { + Type dstMemTyArg, Type /*indicesMemTyArg*/, + Value atomVal, Value src, Value dst, + Value /*indices*/) const { auto srcMemTy = cast(srcMemTyArg); auto dstMemTy = cast(dstMemTyArg); @@ -328,8 +332,9 @@ LogicalResult CopyOpCDNA3BufferCopyLDSType::emitAtomCall(OpBuilder &builder, Loc LogicalResult CopyOpCDNA3BufferCopyLDSType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, - Value atomVal, Value src, Value dst, + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, Value src, + Value dst, Value indices, Value pred) const { OpBuilder::InsertionGuard guard(builder); auto predMemTy = cast(predMemTyArg); @@ -337,7 +342,8 @@ LogicalResult CopyOpCDNA3BufferCopyLDSType::emitAtomCall(OpBuilder &builder, Loc auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } // --- CopyOpCDNA3BufferAtomic --- @@ -489,8 +495,9 @@ FailureOr CopyOpCDNA3BufferAtomicType::emitAtomCallSSA( LogicalResult CopyOpCDNA3BufferAtomicType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, Value src, - Value dst) const { + Type dstMemTyArg, Type /*indicesMemTyArg*/, + Value atomVal, Value src, Value dst, + Value /*indices*/) const { auto srcMemTy = cast(srcMemTyArg); auto srcSSATy = fly::RegMem2SSAType(srcMemTy, /*llvmCompatibleType=*/true); Value srcVal = LLVM::LoadOp::create(builder, loc, srcSSATy, src); @@ -503,15 +510,16 @@ LogicalResult CopyOpCDNA3BufferAtomicType::emitAtomCall(OpBuilder &builder, Loca LogicalResult CopyOpCDNA3BufferAtomicType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, - Value atomVal, Value src, Value dst, - Value pred) const { + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, Value src, + Value dst, Value indices, Value pred) const { auto predMemTy = cast(predMemTyArg); Value predVal = LLVM::LoadOp::create(builder, loc, predMemTy.getElemTy(), pred); auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } } // namespace mlir::fly_rocdl diff --git a/lib/Dialect/FlyROCDL/CDNA4/CopyAtom.cpp b/lib/Dialect/FlyROCDL/CDNA4/CopyAtom.cpp index cdf88650a..856a2334d 100644 --- a/lib/Dialect/FlyROCDL/CDNA4/CopyAtom.cpp +++ b/lib/Dialect/FlyROCDL/CDNA4/CopyAtom.cpp @@ -117,8 +117,9 @@ FailureOr CopyOpCDNA4LdsReadTransposeType::emitAtomCallSSA( LogicalResult CopyOpCDNA4LdsReadTransposeType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, - Value src, Value dst) const { + Type dstMemTyArg, Type /*indicesMemTy*/, + Value atomVal, Value src, Value dst, + Value /*indices*/) const { auto dstSSATy = fly::RegMem2SSAType(cast(dstMemTyArg), true); auto res = emitAtomCallSSA(builder, loc, dstSSATy, copyAtomTyArg, srcMemTyArg, Type{}, atomVal, src, Value{}); @@ -130,8 +131,9 @@ LogicalResult CopyOpCDNA4LdsReadTransposeType::emitAtomCall(OpBuilder &builder, LogicalResult CopyOpCDNA4LdsReadTransposeType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, - Value atomVal, Value src, Value dst, + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, + Value src, Value dst, Value indices, Value pred) const { OpBuilder::InsertionGuard guard(builder); auto predMemTy = cast(predMemTyArg); @@ -139,7 +141,8 @@ LogicalResult CopyOpCDNA4LdsReadTransposeType::emitAtomCall(OpBuilder &builder, auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } LogicalResult CopyOpCDNA4LdsReadTransposeType::verify(function_ref emitError, diff --git a/lib/Dialect/FlyROCDL/GFX1250/CopyAtom.cpp b/lib/Dialect/FlyROCDL/GFX1250/CopyAtom.cpp index 8cba620ea..4fabf27ff 100644 --- a/lib/Dialect/FlyROCDL/GFX1250/CopyAtom.cpp +++ b/lib/Dialect/FlyROCDL/GFX1250/CopyAtom.cpp @@ -215,9 +215,11 @@ Attribute CopyOpGFX1250TDMType::getThrBitLayoutRef() const { LogicalResult CopyOpGFX1250TDMType::verify(function_ref emitError, int32_t rank, int32_t numWarps, int32_t padInterval, int32_t padAmount, int32_t cacheModifier, - bool atomicBarrier, bool earlyTimeout) { + bool atomicBarrier, bool earlyTimeout, TdmMode mode) { if (rank < 1 || rank > static_cast(kMaxTdmRank)) return emitError() << "TDM rank must be in [1, " << kMaxTdmRank << "], got " << rank; + if (mode == TdmMode::Gather && rank != 2) + return emitError() << "TDM gather mode requires rank 2, got " << rank; if (numWarps < 1 || (numWarps & (numWarps - 1)) != 0) return emitError() << "numWarps must be a positive power of two, got " << numWarps; if ((padInterval == 0) != (padAmount == 0)) @@ -239,7 +241,8 @@ FailureOr CopyOpGFX1250TDMType::emitAtomCallSSA(OpBuilder &builder, Locat Type resultTy, Type copyAtomTyArg, Type srcTyArg, Type dstTyArg, Value atomVal, Value src, Value dst) const { - if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, atomVal, src, dst))) + if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, /*indicesMemTy=*/Type{}, + atomVal, src, dst, /*indices=*/Value{}))) return failure(); return Value{}; } @@ -249,16 +252,181 @@ FailureOr CopyOpGFX1250TDMType::emitAtomCallSSA(OpBuilder &builder, Locat Type srcTyArg, Type dstTyArg, Type predTyArg, Value atomVal, Value src, Value dst, Value pred) const { - if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, predTyArg, atomVal, src, - dst, pred))) + if (failed(emitAtomCall(builder, loc, copyAtomTyArg, srcTyArg, dstTyArg, /*indicesMemTy=*/Type{}, + predTyArg, atomVal, src, dst, /*indices=*/Value{}, pred))) return failure(); return Value{}; } +// Gather-mode emit: load a set of rows selected by the `indices` memref operand +// (row-index width from its element type: i16 or 32, count from its static extent) +// into the descriptor index groups, and pack the rank-2 gather descriptor. Base / +// extents / row stride come from atom state (same struct as the tiled path). +static LogicalResult emitTdmGather(OpBuilder &builder, Location loc, CopyOpGFX1250TDMType atomTy, + fly::MemRefType glbMemTy, fly::MemRefType indicesMemTy, + Value indices, Value ldsPtr, bool isLoad, Value atomVal) { + if (!indicesMemTy || !indices) + return mlir::emitError(loc) << "gfx1250 TDM gather requires an `indices` operand"; + + auto layout = dyn_cast(glbMemTy.getLayout()); + if (!layout || !layout.isStaticShape() || layout.rank() != 2) + return failure(); + int32_t outer = layout.getShape().at(0).getLeafAsInt().getValue(); // rows in tile + int32_t rowWidth = layout.getShape().at(1).getLeafAsInt().getValue(); // tile_dim0 + bool hasStaticStride = layout.isStaticStride(); + + int32_t elemBits = glbMemTy.getElemTy().getIntOrFloatBitWidth(); + if (elemBits % 8 != 0) + return failure(); + int32_t elemBytes = elemBits / 8; + if ((elemBytes & (elemBytes - 1)) != 0) + return failure(); + int32_t dataSizeCode = llvm::Log2_32(static_cast(elemBytes)); + + // Row-index operand: width (16/32) from its element type, count from its static + // rank-1 extent (bounded by the descriptor's 8x i32 / 16x i16 index slots). + auto idxLayout = dyn_cast(indicesMemTy.getLayout()); + if (!idxLayout || !idxLayout.isStaticShape() || idxLayout.rank() != 1) + return mlir::emitError(loc) << "gfx1250 TDM gather: indices operand must be a static rank-1 " + "memref of i16 or i32 row indices"; + int32_t indexSize = indicesMemTy.getElemTy().getIntOrFloatBitWidth(); + if (indexSize != 16 && indexSize != 32) + return mlir::emitError(loc) << "gfx1250 TDM gather: indices element type must be i16 or i32"; + int32_t count = idxLayout.getShape().at(0).getLeafAsInt().getValue(); + int32_t maxCount = indexSize == 32 ? 8 : 16; + if (count < 1 || count > maxCount) + return mlir::emitError(loc) << "gfx1250 TDM gather: " << indexSize << "-bit index count must be " + << "in [1, " << maxCount << "], got " << count; + + Type i32Ty = builder.getI32Type(); + Type idxElemTy = builder.getIntegerType(indexSize); + auto idxPtrTy = cast(indices.getType()); + // Load index j from the operand and zero-extend to i32. + auto loadIdx = [&](int32_t j) -> Value { + Value gep = LLVM::GEPOp::create(builder, loc, idxPtrTy, idxElemTy, indices, + ArrayRef{LLVM::GEPArg(j)}); + Value v = LLVM::LoadOp::create(builder, loc, idxElemTy, gep); + if (indexSize == 32) + return v; + return LLVM::ZExtOp::create(builder, loc, i32Ty, v); + }; + + // Pack the loaded indices into 8x i32 descriptor words (32-bit: one per word; + // 16-bit: two per word, lo | hi<<16). Unfilled slots are zero. + Value zeroC = i32Const(builder, loc, 0); + Value c16v = i32Const(builder, loc, 16); + SmallVector words(8, zeroC); + if (indexSize == 32) { + for (int32_t j = 0; j < count; ++j) + words[j] = loadIdx(j); + } else { + for (int32_t w = 0; w < 8; ++w) { + Value lo = (2 * w < count) ? loadIdx(2 * w) : zeroC; + Value hi = (2 * w + 1 < count) ? loadIdx(2 * w + 1) : zeroC; + Value hiSh = arith::ShLIOp::create(builder, loc, hi, c16v); + words[w] = arith::OrIOp::create(builder, loc, lo, hiSh); + } + } + + auto stateField = [&](AtomStateField f) { + return LLVM::ExtractValueOp::create(builder, loc, atomVal, + ArrayRef{*CopyOpGFX1250TDMType::getFieldIndex(f)}); + }; + Value glbBasePtr = stateField(AtomStateField::Base); + Value tensorDim1 = stateField(AtomStateField::Extent0); // rows (OOB on indices) + Value tensorDim0 = stateField(AtomStateField::Extent1); // row width (column OOB) + Value stateStride = stateField(AtomStateField::Stride0); + Value immOffset = stateField(AtomStateField::ImmOffset); + + // Row stride (i64 state; unset sentinel falls back to the static layout stride), + // truncated to the descriptor's 32-bit row-stride slot. + Type i64Ty = builder.getI64Type(); + Value outerStride64; + if (hasStaticStride) { + int64_t s = layout.getStride().at(0).getLeafAsInt().getValue(); + Value unset = + arith::CmpIOp::create(builder, loc, arith::CmpIPredicate::eq, stateStride, + arith::ConstantIntOp::create(builder, loc, kOuterStrideUnset, 64)); + outerStride64 = arith::SelectOp::create( + builder, loc, unset, arith::ConstantIntOp::create(builder, loc, s, 64), stateStride); + } else { + outerStride64 = stateStride; + } + Value outerStride = LLVM::TruncOp::create(builder, loc, i32Ty, outerStride64); + + Value glbBase = LLVM::PtrToIntOp::create(builder, loc, i64Ty, glbBasePtr); + glbBase = arith::AddIOp::create(builder, loc, glbBase, immOffset); + Value ldsAddr = LLVM::PtrToIntOp::create(builder, loc, i32Ty, ldsPtr); + + int32_t padInterval = atomTy.getPadInterval(); + int32_t padAmount = atomTy.getPadAmount(); + FailureOr padOr = computePadEncoding(padInterval, padAmount, elemBits); + if (failed(padOr)) + return mlir::emitError(loc) + << "gfx1250 TDM gather: padding (interval=" << padInterval << ", amount=" << padAmount + << " elements at " << elemBits << "-bit) is not encodable"; + PadEncoding pad = *padOr; + + // GROUP0: pred (gather-index bit [30] set for 32-bit mode, type field [31]), + // lds_addr, glb_lo, glb_hi | type. + int32_t gatherIndexBit = (indexSize == 32) ? 1 : 0; + Value g0s0 = i32Const(builder, loc, 1 | (gatherIndexBit << 30) | (1 << 31)); + Value g0s2 = LLVM::TruncOp::create(builder, loc, i32Ty, glbBase); + Value glbHiRaw = LLVM::LShrOp::create(builder, loc, glbBase, + arith::ConstantIntOp::create(builder, loc, 32, 64)); + Value glbHi = LLVM::TruncOp::create(builder, loc, i32Ty, glbHiRaw); + Value g0s3 = arith::OrIOp::create(builder, loc, glbHi, i32Const(builder, loc, 1 << 31)); + Value dgroup0 = vector::FromElementsOp::create(builder, loc, VectorType::get({4}, i32Ty), + ValueRange{g0s0, ldsAddr, g0s2, g0s3}); + + // GROUP1: config + tensor dims + tile row width + count + row stride. + int32_t g1s0Upper = (dataSizeCode << 16) | ((pad.enable ? 1 : 0) << 20) | (pad.interval << 22) | + (pad.amount << 25); + Value maskRaw = stateField(AtomStateField::WorkgroupMask); + Value maskLow = arith::AndIOp::create(builder, loc, maskRaw, i32Const(builder, loc, 0xFFFF)); + Value g1s0 = arith::OrIOp::create(builder, loc, i32Const(builder, loc, g1s0Upper), maskLow); + + Value mask16 = i32Const(builder, loc, 0xFFFF); + Value td0Lo = arith::AndIOp::create(builder, loc, tensorDim0, mask16); + Value td0Hi = arith::AndIOp::create( + builder, loc, arith::ShRUIOp::create(builder, loc, tensorDim0, c16v), mask16); + Value td1Lo = arith::AndIOp::create(builder, loc, tensorDim1, mask16); + Value td1Hi = arith::AndIOp::create( + builder, loc, arith::ShRUIOp::create(builder, loc, tensorDim1, c16v), mask16); + Value g1s1 = arith::ShLIOp::create(builder, loc, td0Lo, c16v); + Value g1s2 = + arith::OrIOp::create(builder, loc, td0Hi, arith::ShLIOp::create(builder, loc, td1Lo, c16v)); + Value g1s3 = arith::OrIOp::create(builder, loc, td1Hi, i32Const(builder, loc, rowWidth << 16)); + Value g1s4 = i32Const(builder, loc, count & 0xFFFF); // gather tile_dim1 = valid index count + Value dgroup1 = vector::FromElementsOp::create( + builder, loc, VectorType::get({8}, i32Ty), + ValueRange{g1s0, g1s1, g1s2, g1s3, g1s4, outerStride, zeroC, zeroC}); + + // GROUP2 / GROUP3: the packed row-index words. + Value dg2 = vector::FromElementsOp::create(builder, loc, VectorType::get({4}, i32Ty), + ValueRange{words[0], words[1], words[2], words[3]}); + Value dg3 = vector::FromElementsOp::create(builder, loc, VectorType::get({4}, i32Ty), + ValueRange{words[4], words[5], words[6], words[7]}); + Value dg4 = vector::FromElementsOp::create( + builder, loc, VectorType::get({8}, i32Ty), + ValueRange{zeroC, zeroC, zeroC, zeroC, zeroC, zeroC, zeroC, zeroC}); + + uint32_t cachePolicy = static_cast(atomTy.getCacheModifier()); + ArrayAttr noAliasScopes; + if (isLoad) + ROCDL::TensorLoadToLDSOp::create(builder, loc, dgroup0, dgroup1, dg2, dg3, dg4, cachePolicy, + noAliasScopes, noAliasScopes, noAliasScopes); + else + ROCDL::TensorStoreFromLDSOp::create(builder, loc, dgroup0, dgroup1, dg2, dg3, dg4, cachePolicy, + noAliasScopes, noAliasScopes, noAliasScopes); + return success(); +} + LogicalResult CopyOpGFX1250TDMType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Value atomVal, Value src, - Value dst) const { + Type dstMemTyArg, Type indicesMemTyArg, + Value atomVal, Value src, Value dst, + Value indices) const { auto srcMemTy = dyn_cast(srcMemTyArg); auto dstMemTy = dyn_cast(dstMemTyArg); if (!srcMemTy || !dstMemTy) @@ -279,6 +447,12 @@ LogicalResult CopyOpGFX1250TDMType::emitAtomCall(OpBuilder &builder, Location lo // the operand layout (tensor dim order: index 0 = outermost, rank-1 = innermost). fly::MemRefType glbMemTy = isLoad ? srcMemTy : dstMemTy; Value ldsPtr = isLoad ? dst : src; + + if (getMode() == TdmMode::Gather) + return emitTdmGather(builder, loc, *this, glbMemTy, + dyn_cast_or_null(indicesMemTyArg), indices, ldsPtr, isLoad, + atomVal); + int32_t rank = getRank(); auto layout = dyn_cast(glbMemTy.getLayout()); @@ -524,14 +698,16 @@ LogicalResult CopyOpGFX1250TDMType::emitAtomCall(OpBuilder &builder, Location lo LogicalResult CopyOpGFX1250TDMType::emitAtomCall(OpBuilder &builder, Location loc, Type copyAtomTyArg, Type srcMemTyArg, - Type dstMemTyArg, Type predMemTyArg, Value atomVal, - Value src, Value dst, Value pred) const { + Type dstMemTyArg, Type indicesMemTyArg, + Type predMemTyArg, Value atomVal, Value src, + Value dst, Value indices, Value pred) const { OpBuilder::InsertionGuard guard(builder); auto predMemTy = cast(predMemTyArg); Value predVal = LLVM::LoadOp::create(builder, loc, predMemTy.getElemTy(), pred); auto ifOp = scf::IfOp::create(builder, loc, TypeRange{}, predVal, /*withElse=*/false); builder.setInsertionPointToStart(&ifOp.getThenRegion().front()); - return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, atomVal, src, dst); + return emitAtomCall(builder, loc, copyAtomTyArg, srcMemTyArg, dstMemTyArg, indicesMemTyArg, + atomVal, src, dst, indices); } } // namespace mlir::fly_rocdl diff --git a/python/flydsl/expr/derived.py b/python/flydsl/expr/derived.py index 9e81b3891..7cdebb490 100644 --- a/python/flydsl/expr/derived.py +++ b/python/flydsl/expr/derived.py @@ -221,7 +221,17 @@ def gather(copy_atom, base_iter, offset_tensor, dst_tensor, *, pred=None): ``base_iter``. The reconstructed source view uses the copy atom's source value layout, while ``dst_tensor[(None, v), rest]`` supplies the matching destination ``(AtomV,)`` slice. + + Whole-tile atoms (``copy_atom.copy_rank > 1``, e.g. the gfx1250 TDM gather + atom) instead move all rows in one hardware call: ``offset_tensor`` is the + row-index operand, ``dst_tensor`` the destination tile, and ``base_iter`` only + supplies the global shape/direction token (its pointer is unused — the base is + atom state). """ + if getattr(copy_atom, "copy_rank", 1) > 1: + src_token = make_view(base_iter, get_layout(dst_tensor)) + copy_atom_call(copy_atom, src_token, dst_tensor, indices=offset_tensor, pred=pred) + return src_layout = copy_atom.layout_src_tv[1] for off, dst_v, pred_v in _gather_scatter_expand(offset_tensor, dst_tensor, pred): src_v = make_view(base_iter + off, src_layout) @@ -244,7 +254,17 @@ def scatter(copy_atom, src_tensor, base_iter, offset_tensor, *, pred=None): For each ``(v, rest)`` instance, ``src_tensor[(None, v), rest]`` supplies the source ``(AtomV,)`` slice. The reconstructed destination view uses the copy atom's destination value layout at ``base_iter + offset_tensor[v, rest]``. + + Whole-tile atoms (``copy_atom.copy_rank > 1``, e.g. the gfx1250 TDM scatter + atom) instead move all rows in one hardware call: ``offset_tensor`` is the + row-index operand, ``src_tensor`` the source tile, and ``base_iter`` only + supplies the global shape/direction token (its pointer is unused — the base is + atom state). """ + if getattr(copy_atom, "copy_rank", 1) > 1: + dst_token = make_view(base_iter, get_layout(src_tensor)) + copy_atom_call(copy_atom, src_tensor, dst_token, indices=offset_tensor, pred=pred) + return dst_layout = copy_atom.layout_dst_tv[1] for off, src_v, pred_v in _gather_scatter_expand(offset_tensor, src_tensor, pred): dst_v = make_view(base_iter + off, dst_layout) diff --git a/python/flydsl/expr/primitive.py b/python/flydsl/expr/primitive.py index d720dad19..491ec824f 100644 --- a/python/flydsl/expr/primitive.py +++ b/python/flydsl/expr/primitive.py @@ -1045,8 +1045,8 @@ def atom_set_value(atom, field, value): @dsl_loc_tracing -def copy_atom_call(copy_atom, src, dst, *, pred=None): - return fly.copy_atom_call(copy_atom, src, dst, pred=pred) +def copy_atom_call(copy_atom, src, dst, *, indices=None, pred=None): + return fly.copy_atom_call(copy_atom, src, dst, indices=indices, pred=pred) @dsl_loc_tracing @@ -1103,8 +1103,8 @@ def mma_make_fragment(operand_id, tiled_mma, input, *, stages=None): @dsl_loc_tracing -def copy(copy_atom, src, dst, *, pred=None, **kwargs): - return fly.copy(copy_atom.set_value(kwargs), src, dst, pred=pred) +def copy(copy_atom, src, dst, *, indices=None, pred=None, **kwargs): + return fly.copy(copy_atom.set_value(kwargs), src, dst, indices=indices, pred=pred) @dsl_loc_tracing diff --git a/python/flydsl/expr/rocdl/universal.py b/python/flydsl/expr/rocdl/universal.py index 8d30a922c..da850a123 100644 --- a/python/flydsl/expr/rocdl/universal.py +++ b/python/flydsl/expr/rocdl/universal.py @@ -178,6 +178,11 @@ def WMMAScale( ) +# gfx1250 TDM copy modes (must match FlyROCDL_TdmMode in Atom.td). +TDM_MODE_TILED = 0 +TDM_MODE_GATHER = 1 + + def TDM( rank, num_warps, @@ -186,6 +191,7 @@ def TDM( cache_modifier=0, atomic_barrier=False, early_timeout=False, + mode=TDM_MODE_TILED, ): """Create a gfx1250 N-D TDM (Tensor Data Mover) Global<->LDS copy atom *type*. @@ -197,6 +203,11 @@ def TDM( ``atomic_barrier`` (descriptor bit 18, HW auto-barrier) and ``early_timeout`` (bit 21, multicast-load GL1 knob) set compile-time descriptor config bits. + ``mode`` is :data:`TDM_MODE_TILED` (contiguous N-D tile) or + :data:`TDM_MODE_GATHER` (rank-2 row gather/scatter; the row indices are passed + as the ``indices`` operand of ``fx.copy_atom_call`` / via ``fx.gather`` / + ``fx.scatter``, and the index width is taken from that operand's element type). + The tile descriptor (global base pointer, per-dim extent for out-of-bounds handling, per-dim stride) plus the MCAST ``workgroup_mask`` are runtime atom state set via ``fx.atom.set_value``. :func:`make_tdm_atom` builds the atom and @@ -210,6 +221,7 @@ def TDM( cache_modifier, atomic_barrier=atomic_barrier, early_timeout=early_timeout, + mode=mode, ) @@ -280,6 +292,7 @@ def make_tdm_atom( cache_modifier=0, atomic_barrier=False, early_timeout=False, + mode=TDM_MODE_TILED, ) -> object: """Build a gfx1250 N-D TDM copy atom carrying ``tensor``'s tile descriptor. @@ -324,6 +337,7 @@ def make_tdm_atom( cache_modifier, atomic_barrier=atomic_barrier, early_timeout=early_timeout, + mode=mode, ) atom = make_copy_atom(copy_op, tensor.element_type) atom = atom_set_value(atom, "base", get_iter(tensor)) @@ -344,6 +358,67 @@ def make_tdm_atom( return atom +def make_tdm_gather_atom( + tensor: Tensor, + *, + num_rows=None, + row_width_bound=None, + row_stride=None, + num_warps=1, + pad_interval=0, + pad_amount=0, + cache_modifier=0, + workgroup_mask=0, +) -> object: + """Build a gfx1250 TDM *gather* copy atom (mode = gather, rank 2). + + Gathers a set of rows selected by an explicit row-index operand: issue the copy + with ``fx.gather(atom, base_iter, row_indices, dst)`` / ``fx.scatter`` (or + ``fx.copy_atom_call(atom, global_tile, lds, indices=row_indices)``). The row + indices ride as the ``indices`` operand — their element type (i16 / i32) sets the + index width and their count sets the number of rows gathered; they are NOT atom + state. + + The atom carries only the global descriptor as runtime state: ``base`` pointer, + OOB extents (``extent_0`` = row count / ``num_rows``, ``extent_1`` = row width / + ``row_width_bound``; ``None`` leaves an axis unclamped) and the row stride + (``extent`` order matches the rank-2 gather tile: dim0 = rows, dim1 = row width). + ``row_stride`` (elements) overrides dim0 stride; ``None`` falls back to the tile + memref's static layout stride. + """ + from ..primitive import atom_set_value, make_copy_atom + + NO_CLAMP = 0x7FFFFFFF + STRIDE_UNSET = -0x80000000 # matches kOuterStrideUnset in CopyAtom.cpp + + def _i32(v): + return v if isinstance(v, Int32) else Int32(v) + + copy_op = CopyOpGFX1250TDMType.get( + 2, + num_warps, + pad_interval, + pad_amount, + cache_modifier, + atomic_barrier=False, + early_timeout=False, + mode=TDM_MODE_GATHER, + ) + atom = make_copy_atom(copy_op, tensor.element_type) + atom = atom_set_value(atom, "base", get_iter(tensor)) + atom = atom_set_value(atom, "extent_0", Int32(NO_CLAMP) if num_rows is None else _i32(num_rows)) + atom = atom_set_value(atom, "extent_1", Int32(NO_CLAMP) if row_width_bound is None else _i32(row_width_bound)) + stride = ( + Int64(STRIDE_UNSET) + if row_stride is None + else (row_stride if isinstance(row_stride, Int64) else Int64(row_stride)) + ) + atom = atom_set_value(atom, "stride_0", stride) + if workgroup_mask != 0: + atom = atom_set_value(atom, "workgroup_mask", _i32(workgroup_mask)) + return atom + + def advance_tdm_atom(atom, byte_offset) -> object: """Return a TDM atom with its global byte offset (``imm_offset``) set. diff --git a/python/flydsl/expr/typing.py b/python/flydsl/expr/typing.py index b0cd30ade..4d085080d 100644 --- a/python/flydsl/expr/typing.py +++ b/python/flydsl/expr/typing.py @@ -1080,6 +1080,12 @@ class CopyAtom(BuiltinDslType): def val_bits(self): return self.type.val_bits + @property + def copy_rank(self): + """Memref dims moved per copy_atom_call (whole-tile atoms like TDM return + N; 1 otherwise). ``fx.gather`` / ``fx.scatter`` route on this.""" + return self.type.copy_rank + @property def thr_layout(self): return static(self.type.thr_layout) diff --git a/tests/mlir/Conversion/gfx1250_atoms_neg.mlir b/tests/mlir/Conversion/gfx1250_atoms_neg.mlir index 760ad3a60..32dab6aad 100644 --- a/tests/mlir/Conversion/gfx1250_atoms_neg.mlir +++ b/tests/mlir/Conversion/gfx1250_atoms_neg.mlir @@ -34,7 +34,7 @@ func.func @bad_mma_blocksize( // CHECK: numWarps must be a positive power of two, got 3 func.func @bad_tdm_warps( - %a: !fly.copy_atom, 0>) { + %a: !fly.copy_atom, 0>) { return } @@ -42,7 +42,7 @@ func.func @bad_tdm_warps( // CHECK: padInterval and padAmount must both be zero or both non-zero func.func @bad_tdm_pad( - %a: !fly.copy_atom, 0>) { + %a: !fly.copy_atom, 0>) { return } @@ -52,7 +52,7 @@ func.func @bad_tdm_pad( // interval -> a wrong encoded bitfield). Caught statically by the verifier. // CHECK: padInterval must be a power of two (in elements), got 48 func.func @bad_tdm_pad_pow2( - %a: !fly.copy_atom, 0>) { + %a: !fly.copy_atom, 0>) { return } diff --git a/tests/mlir/Conversion/tdm_gather_gfx1250.mlir b/tests/mlir/Conversion/tdm_gather_gfx1250.mlir new file mode 100644 index 000000000..91fa7c9d3 --- /dev/null +++ b/tests/mlir/Conversion/tdm_gather_gfx1250.mlir @@ -0,0 +1,72 @@ +// SPDX-License-Identifier: Apache-2.0 +// Copyright (c) 2026 FlyDSL Project Contributors +// RUN: %fly-opt %s --convert-fly-to-rocdl | FileCheck %s + +// gfx1250 TDM gather (mode = gather): rows selected by an `indices` memref operand +// on copy_atom_call are moved Global<->LDS. The row-index width is taken from the +// operand element type (i32 or i16) and the count from its static rank-1 extent; +// the indices are loaded and packed into descriptor groups 2/3. Base / extents / +// row stride come from atom state (same struct as the tiled TDM atom). +// Global -> Shared => rocdl.tensor.load.to.lds +// Shared -> Global => rocdl.tensor.store.from.lds + +// ----- + +// 32-bit gather load: 8 i32 indices are GEP+loaded from the operand and placed one +// per descriptor word. GROUP0 pred = 1 | (1<<30 gather-index) | (1<<31 type) = +// 0xC0000001 = -1073741823. row_width (tile inner = 64) packs into GROUP1 s3 at bit +// 16 (64<<16 = 4194304); gather tile_dim1 = index count = 8. + +// CHECK-LABEL: @gather_load_i32 +func.func @gather_load_i32( + %atom: !fly.copy_atom, 0>, + %nrows: i32, + %src: !fly.memref, + %dst: !fly.memref, + %idx: !fly.memref) { + // extent_0 (slot 2) = tensor_dim1 = runtime row count for OOB on the indices. + %a1 = fly.atom.set_value(%atom, "extent_0", %nrows) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> + // Indices are GEP+loaded from the operand and packed into the descriptor. + // CHECK: llvm.getelementptr + // CHECK: llvm.load + // pred word with the gather-index + type bits set. + // CHECK-DAG: arith.constant -1073741823 : i32 + // row_width (64) packed at bit 16. + // CHECK-DAG: arith.constant 4194304 : i32 + // CHECK: rocdl.tensor.load.to.lds %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}} cachepolicy 0 : vector<4xi32>, vector<8xi32> + fly.copy_atom_call(%a1, %src, %dst indices = %idx) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref, !fly.memref) -> () + return +} + +// ----- + +// Store direction (Shared -> Global) -> tensor.store.from.lds. +// CHECK-LABEL: @gather_store_i32 +func.func @gather_store_i32( + %atom: !fly.copy_atom, 0>, + %src: !fly.memref, + %dst: !fly.memref, + %idx: !fly.memref) { + // CHECK: rocdl.tensor.store.from.lds %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}} cachepolicy 0 : vector<4xi32>, vector<8xi32> + fly.copy_atom_call(%atom, %src, %dst indices = %idx) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref, !fly.memref) -> () + return +} + +// ----- + +// 16-bit gather load: up to 16 i16 indices, packed two per descriptor word +// (lo | hi<<16). GROUP0 pred = 1 | (1<<31 type) = 0x80000001 = -2147483647 (the +// gather-index bit [30] is NOT set for 16-bit mode). +// CHECK-LABEL: @gather_load_i16 +func.func @gather_load_i16( + %atom: !fly.copy_atom, 0>, + %src: !fly.memref, + %dst: !fly.memref, + %idx: !fly.memref) { + // 16-bit indices are zero-extended and packed two per word. + // CHECK-DAG: llvm.zext + // CHECK-DAG: arith.constant -2147483647 : i32 + // CHECK: rocdl.tensor.load.to.lds %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}} cachepolicy 0 : vector<4xi32>, vector<8xi32> + fly.copy_atom_call(%atom, %src, %dst indices = %idx) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref, !fly.memref) -> () + return +} diff --git a/tests/mlir/Conversion/tdm_gfx1250.mlir b/tests/mlir/Conversion/tdm_gfx1250.mlir index bd81a6f60..cfd922dc7 100644 --- a/tests/mlir/Conversion/tdm_gfx1250.mlir +++ b/tests/mlir/Conversion/tdm_gfx1250.mlir @@ -15,7 +15,7 @@ // CHECK-LABEL: @test_tdm_type // CHECK-SAME: (%{{.*}}: !llvm.struct<(i32, ptr<1>, i32, i32, i32, i32, i32, i64, i64, i64, i64, i64)>) func.func @test_tdm_type( - %atom: !fly.copy_atom, 0>) { + %atom: !fly.copy_atom, 0>) { return } @@ -27,13 +27,13 @@ func.func @test_tdm_type( // CHECK-LABEL: @test_tdm_load func.func @test_tdm_load( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %oe: i32, %ie: i32, %os: i64, %src: !fly.memref, %dst: !fly.memref) { - %a1 = fly.atom.set_value(%atom, "extent_0", %oe) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> - %a2 = fly.atom.set_value(%a1, "extent_1", %ie) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> - %a3 = fly.atom.set_value(%a2, "stride_0", %os) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> + %a1 = fly.atom.set_value(%atom, "extent_0", %oe) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> + %a2 = fly.atom.set_value(%a1, "extent_1", %ie) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> + %a3 = fly.atom.set_value(%a2, "stride_0", %os) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> // stride_0 fallback: select(stride == unset-sentinel (i64), static_layout_stride, stride) // CHECK-DAG: %[[STRIDE:.*]] = llvm.extractvalue %{{.*}}[7] : !llvm.struct<(i32, ptr<1>, i32, i32, i32, i32, i32, i64, i64, i64, i64, i64)> // CHECK-DAG: %[[SENT:.*]] = arith.constant -2147483648 : i64 @@ -46,7 +46,7 @@ func.func @test_tdm_load( // CHECK: arith.subi // CHECK: arith.maxsi // CHECK: rocdl.tensor.load.to.lds %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}} cachepolicy 0 : vector<4xi32>, vector<8xi32> - fly.copy_atom_call(%a3, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%a3, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -56,11 +56,11 @@ func.func @test_tdm_load( // CHECK-LABEL: @test_tdm_store func.func @test_tdm_store( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %src: !fly.memref, %dst: !fly.memref) { // CHECK: rocdl.tensor.store.from.lds %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}}, %{{.*}} cachepolicy 0 : vector<4xi32>, vector<8xi32> - fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -71,14 +71,14 @@ func.func @test_tdm_store( // CHECK-LABEL: @test_tdm_load_warps func.func @test_tdm_load_warps( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %src: !fly.memref, %dst: !fly.memref) { // CHECK: %[[WID:.*]] = rocdl.wave.id : i32 // CHECK-DAG: arith.remui %[[WID]] // CHECK-DAG: arith.divui %[[WID]] // CHECK: rocdl.tensor.load.to.lds - fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -89,15 +89,15 @@ func.func @test_tdm_load_warps( // CHECK-LABEL: @test_tdm_load_3d func.func @test_tdm_load_3d( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %s0: i64, %s1: i64, %src: !fly.memref, %dst: !fly.memref) { - %a1 = fly.atom.set_value(%atom, "stride_0", %s0) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> - %a2 = fly.atom.set_value(%a1, "stride_1", %s1) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> + %a1 = fly.atom.set_value(%atom, "stride_0", %s0) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> + %a2 = fly.atom.set_value(%a1, "stride_1", %s1) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> // CHECK: llvm.extractvalue %{{.*}}[8] : !llvm.struct<(i32, ptr<1>, i32, i32, i32, i32, i32, i64, i64, i64, i64, i64)> // CHECK: rocdl.tensor.load.to.lds - fly.copy_atom_call(%a2, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%a2, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -110,12 +110,12 @@ func.func @test_tdm_load_3d( // CHECK-LABEL: @test_tdm_load_pad func.func @test_tdm_load_pad( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %src: !fly.memref, %dst: !fly.memref) { // CHECK-DAG: arith.constant 118554624 : i32 // CHECK: rocdl.tensor.load.to.lds - fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -128,14 +128,14 @@ func.func @test_tdm_load_pad( // CHECK-LABEL: @test_tdm_store_pad func.func @test_tdm_store_pad( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %src: !fly.memref, %dst: !fly.memref) { // CHECK-DAG: arith.constant 118554624 : i32 // tile_dim0 stays 64 (0x40) -> 64 << 16 = 4194304; NOT 72 (64+pad 8). // CHECK-DAG: arith.constant 4194304 : i32 // CHECK: rocdl.tensor.store.from.lds - fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -146,18 +146,18 @@ func.func @test_tdm_store_pad( // CHECK-LABEL: @test_tdm_load_mcast func.func @test_tdm_load_mcast( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %mask: i32, %src: !fly.memref, %dst: !fly.memref) { // CHECK: %[[A1:.*]] = llvm.insertvalue %{{.*}}, %{{.*}}[0] - %a1 = fly.atom.set_value(%atom, "workgroup_mask", %mask) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> + %a1 = fly.atom.set_value(%atom, "workgroup_mask", %mask) : (!fly.copy_atom, 0>, i32) -> !fly.copy_atom, 0> // CHECK: %[[M:.*]] = llvm.extractvalue %[[A1]][0] // CHECK-DAG: %[[MLOW:.*]] = arith.andi %[[M]], %{{.*}} // CHECK-DAG: %[[UPPER:.*]] = arith.constant 65536 : i32 // CHECK: arith.ori %[[UPPER]], %[[MLOW]] // CHECK: rocdl.tensor.load.to.lds - fly.copy_atom_call(%a1, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%a1, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -169,12 +169,12 @@ func.func @test_tdm_load_mcast( // CHECK-LABEL: @test_tdm_load_barrier_timeout func.func @test_tdm_load_barrier_timeout( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %src: !fly.memref, %dst: !fly.memref) { // CHECK-DAG: arith.constant 2424832 : i32 // CHECK: rocdl.tensor.load.to.lds - fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } @@ -185,17 +185,17 @@ func.func @test_tdm_load_barrier_timeout( // CHECK-LABEL: @test_tdm_load_imm_offset func.func @test_tdm_load_imm_offset( - %atom: !fly.copy_atom, 0>, + %atom: !fly.copy_atom, 0>, %off: i64, %src: !fly.memref, %dst: !fly.memref) { // CHECK: %[[A1:.*]] = llvm.insertvalue %{{.*}}, %{{.*}}[11] : !llvm.struct<(i32, ptr<1>, i32, i32, i32, i32, i32, i64, i64, i64, i64, i64)> - %a1 = fly.atom.set_value(%atom, "imm_offset", %off) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> + %a1 = fly.atom.set_value(%atom, "imm_offset", %off) : (!fly.copy_atom, 0>, i64) -> !fly.copy_atom, 0> // CHECK-DAG: %[[BASE:.*]] = llvm.extractvalue %[[A1]][1] // CHECK-DAG: %[[IMM:.*]] = llvm.extractvalue %[[A1]][11] // CHECK-DAG: %[[BI:.*]] = llvm.ptrtoint %[[BASE]] : !llvm.ptr<1> to i64 // CHECK: arith.addi %[[BI]], %[[IMM]] : i64 // CHECK: rocdl.tensor.load.to.lds - fly.copy_atom_call(%a1, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + fly.copy_atom_call(%a1, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } diff --git a/tests/mlir/Conversion/tdm_gfx1250_neg.mlir b/tests/mlir/Conversion/tdm_gfx1250_neg.mlir index 1ac92fd52..4d24d0189 100644 --- a/tests/mlir/Conversion/tdm_gfx1250_neg.mlir +++ b/tests/mlir/Conversion/tdm_gfx1250_neg.mlir @@ -16,7 +16,7 @@ func.func @bad_tdm_pad_interval_dw( %src: !fly.memref, %dst: !fly.memref) { - %atom = fly.make_copy_atom {valBits = 0 : i32} : !fly.copy_atom, 0> - fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + %atom = fly.make_copy_atom {valBits = 0 : i32} : !fly.copy_atom, 0> + fly.copy_atom_call(%atom, %src, %dst) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } diff --git a/tests/mlir/Transforms/expand_copy_tdm_wholetile.mlir b/tests/mlir/Transforms/expand_copy_tdm_wholetile.mlir index 2b35bcdca..bfbd17546 100644 --- a/tests/mlir/Transforms/expand_copy_tdm_wholetile.mlir +++ b/tests/mlir/Transforms/expand_copy_tdm_wholetile.mlir @@ -15,7 +15,7 @@ func.func @tdm_wholetile( %g: !fly.memref, %s: !fly.memref) { - %atom = fly.make_copy_atom {valBits = 0 : i32} : !fly.copy_atom, 0> - fly.copy(%atom, %g, %s) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () + %atom = fly.make_copy_atom {valBits = 0 : i32} : !fly.copy_atom, 0> + fly.copy(%atom, %g, %s) : (!fly.copy_atom, 0>, !fly.memref, !fly.memref) -> () return } diff --git a/tests/unit/test_gfx1250_atoms.py b/tests/unit/test_gfx1250_atoms.py index 186e03b68..183c29067 100644 --- a/tests/unit/test_gfx1250_atoms.py +++ b/tests/unit/test_gfx1250_atoms.py @@ -69,8 +69,9 @@ def test_tdm2d_type_roundtrip(): t = U.TDM(2, 1) assert "gfx1250.tdm<" in str(t) assert "rank = 2" in str(t) - # Defaults: barrier / timeout = false. + # Defaults: barrier / timeout = false, mode = tiled. assert "barrier = false, timeout = false" in str(t) + assert "mode = tiled" in str(t) assert ir.Type.parse(str(t)) == t t2 = U.TDM(3, 8, pad_interval=64, pad_amount=8, cache_modifier=2) @@ -83,3 +84,26 @@ def test_tdm2d_type_roundtrip(): t3 = U.TDM(1, 1, atomic_barrier=True, early_timeout=True) assert "barrier = true, timeout = true" in str(t3) assert ir.Type.parse(str(t3)) == t3 + + +def test_tdm_gather_mode_roundtrip(): + with _ctx(), ir.Location.unknown(): + from flydsl._mlir.dialects import fly # noqa: F401 + from flydsl._mlir.dialects import fly_rocdl # noqa: F401 + from flydsl._mlir.dialects.fly import CopyAtomType + from flydsl.expr.rocdl import universal as U + + assert U.TDM_MODE_TILED == 0 and U.TDM_MODE_GATHER == 1 + + g = U.TDM(2, 1, mode=U.TDM_MODE_GATHER) + assert "mode = gather" in str(g) + assert ir.Type.parse(str(g)) == g + + # A whole-tile atom reports copy_rank == rank (2), which the DSL + # gather/scatter routing keys on. + cat = CopyAtomType.get(copy_op=g, val_bits=16) + assert cat.copy_rank == 2 + + # gather mode requires rank 2. + with pytest.raises(Exception): + ir.Type.parse(str(U.TDM(3, 1)).replace("mode = tiled", "mode = gather"))