From 019830335862e3592eb1ad62e3301c979c056a85 Mon Sep 17 00:00:00 2001 From: Zheng QiHang Date: Mon, 21 Sep 2026 16:35:37 +0800 Subject: [PATCH] feat(shard): express swizzled composed layouts as CuTe does Port CuTe's Swizzle as a frozen value type that ComposedLayout may carry in `inner`, so a shared-memory tile can state the XOR swizzle a WGMMA operand needs instead of only an affine stride rule. The layout algebra gains CuTe's swizzle specializations -- apply, cosize, coalesce, the involution inverses and the composition that canonicalizes a Swizzle back to the inner side rather than into an `outer` that has no domain. The CUDA path emits the composition unchanged, as cute::make_composed_layout(cute::Swizzle{}, cute::Int{}, ...), for both the layout type and the layout value. Because CuTe deletes stride() on a ComposedLayout, the device runtime now reads steps off the affine part of a shard layout, and projects a composed one by folding the instance's origin into the composition's offset instead of advancing the engine pointer: the swizzle runs on the whole tensor's index, so pointer arithmetic would hand two instances one address. A Swizzle states where an element lives, never which element a logical coordinate names, so shape, domain rank, the axis numbering a Split references and the access relation an op states are all read off `outer` and are the unswizzled ones. Transpose, Reshape and Slice therefore carry the same Swizzle through, a window's start being a constant shift of the index that `offset` already holds. A register engine refuses it by name rather than rebuilding the layout from register strides. --- docs/spec/runtime.md | 5 + docs/spec/shard.md | 64 ++++++- .../runtime/cuda/layout/shard_layout.cuh | 25 ++- .../runtime/cuda/tensor_view/shard_tensor.cuh | 28 ++- src/tilefoundry/__init__.py | 3 +- .../codegen/cuda/tir/memory/tensor_view.py | 79 ++++++-- src/tilefoundry/inspection/printer_base.py | 7 +- src/tilefoundry/ir/hir/tensor/slice.py | 37 +++- src/tilefoundry/ir/types/shard/__init__.py | 13 +- src/tilefoundry/ir/types/shard/layout.py | 73 +++++++- .../ir/types/shard/layout_algebra.py | 177 ++++++++++++++++-- tests/ops/tir/cuda/test_swizzle.py | 124 ++++++++++++ 12 files changed, 587 insertions(+), 48 deletions(-) create mode 100644 tests/ops/tir/cuda/test_swizzle.py diff --git a/docs/spec/runtime.md b/docs/spec/runtime.md index 3c286a39..e46738dc 100644 --- a/docs/spec/runtime.md +++ b/docs/spec/runtime.md @@ -1044,6 +1044,9 @@ template concept ShardTensorLike = detail::is_shard_tensor>::value; +template +inline constexpr bool shard_layout_is_composed_v; + template CUTE_HOST_DEVICE decltype(auto) local_tensor(T &&t); template @@ -1058,6 +1061,8 @@ template CUTE_HOST_DEVICE constexpr int shard_mesh_instances(); - constraints: - `engine` is a CuTe tensor or view, never a raw pointer: residency lives on the engine type, and `data()` drops it. + - A shard layout whose layout is a `cute::ComposedLayout` -- which is how a [`Swizzle`](./shard.md#41-swizzle) reaches the device -- is projected with its function kept: `shard_layout_is_composed_v` picks that path, and `local_tensor` returns a tensor over the engine's own pointer whose layout is the same composition, this instance's origin folded into its offset. Advancing the pointer instead would run the function on a slice-relative index and hand two instances one address, so the offset MUST NOT move to the engine. + - A composed layout takes that path even when no mesh axis splits it. The whole-tensor shortcut returns the backing engine, whose layout states no function at all, so taking it would drop the swizzle. - `local_tensor`, `local_view_t` and `ShardTensorLike` are public: a step every op must take, and the word an op writes its own constraints in, are not implementation details. `is_shard_tensor` -- how `ShardTensorLike` is decided, which nothing outside needs -- stays in `detail`. - A tensor whose layout is a `ShardLayout` has distributed semantics -- it is the whole tensor, and each instance owns the slice its shard layout gives it. - One attr per mesh axis, against the mesh's axes flattened: a mesh naming several levels states them grouped one nest per level, and its rank is then how many levels it names rather than how many axes the attrs answer for. diff --git a/docs/spec/shard.md b/docs/spec/shard.md index 7decb288..8ab3d31b 100644 --- a/docs/spec/shard.md +++ b/docs/spec/shard.md @@ -178,7 +178,7 @@ class ComposedLayout(LayoutBase): outer: attribute; Domain-side layout applied first, or identity. """ - inner: LayoutBase | None + inner: LayoutBase | Swizzle | None offset: int outer: LayoutBase | None ``` @@ -188,8 +188,13 @@ class ComposedLayout(LayoutBase): `apply(None, value) == value`. - `shape` and `domain_rank` come from `outer` when it is present, otherwise from `inner`; a composition with no explicit domain has `shape == ()`. - - either non-`None` component MAY be any `LayoutBase`, including a - `ShardLayout` that preserves an earlier distribution. + A `Swizzle` states no domain, so `inner=Swizzle` with `outer=None` has + `shape == ()` and is not a complete tensor layout. + - `outer`, and an `inner` that is a `LayoutBase`, MAY be any `LayoutBase`, + including a `ShardLayout` that preserves an earlier distribution. + - `inner` MAY instead be a [`Swizzle`](#41-swizzle), which is the only + non-affine mapping this IR states. `outer` MUST NOT be one: a `Swizzle` + has no domain to be the domain-side component of. Field meanings: @@ -206,6 +211,59 @@ domain, so the composition inherits them from `inner`. Therefore, when an outer `ShardLayout` binds a `ComposedLayout`, a `Split(k)` attr still references the composition's stable `shape` / `domain_rank` contract. +### 4.1 `Swizzle` + +Mirrors CuTe `Swizzle`: the XOR offset functor a swizzled +shared-memory layout applies to its index. + +```python +class Swizzle: + """Describe the XOR permutation CuTe `Swizzle` applies to an index. + + Attributes: + bits: attribute; CuTe `BBits`, how many bits are XORed. + base: attribute; CuTe `MBase`, how many low bits are left alone. + shift: attribute; CuTe `SShift`, the signed source-to-target distance. + """ + + bits: int + base: int + shift: int + + def __call__(self, offset: int) -> int: ... +``` + +**Terms.** The *Y bits* are the ones read out of the index +(`yyy_msk = ((1 << bits) - 1) << (base + max(0, shift))`); the *Z bits* are +the ones they are XORed onto +(`zzz_msk = ((1 << bits) - 1) << (base - min(0, shift))`). + +- constraints: + - `bits >= 0`, `base >= 0`, and `abs(shift) >= bits`, which is what keeps + the Y and Z ranges from overlapping. + - `swizzle(offset) == offset ^ shiftr(offset & yyy_msk, shift)`, the CuTe + formula unchanged. + - Because the two ranges do not overlap, a `Swizzle` is an involution: + `swizzle(swizzle(offset)) == offset`, so it is its own left and right + inverse. `bits == 0` is the identity. + - A `Swizzle` is a mapping on an index, not a layout. It is not a + `LayoutBase`, it states no `shape`, and it MUST NOT be a + `TensorType.layout` or a `ShardLayout.layout` on its own. It reaches a + tensor only as `ComposedLayout.inner`. + - A `Swizzle` permutes bits inside the codomain `outer` already spans, so + `cosize(ComposedLayout(Swizzle, offset, outer)) == cosize(outer)` and the + backing allocation is unchanged by it. + - A `Swizzle` states where an element lives, never which element a logical + coordinate names. `shape`, `domain_rank`, the axis numbering a `Split` + references, and the access relation an op states over a value are all read + off `outer` and MUST be the same as the unswizzled layout's. This is why an + op that only relabels the domain -- `Transpose`, `Reshape`, `Slice` -- + carries the same `Swizzle` through: a window's start is a constant shift of + the index, which `offset` already states. + - A swizzled `ComposedLayout` MUST NOT serve as a `Mesh` execution scope: + [§9](#9-layout-construction-and-mesh-scope-projection) admits an identity + `inner` only. + --- ## 5. `Mesh` diff --git a/include/tilefoundry/runtime/cuda/layout/shard_layout.cuh b/include/tilefoundry/runtime/cuda/layout/shard_layout.cuh index 1db06b38..7dbb83a8 100644 --- a/include/tilefoundry/runtime/cuda/layout/shard_layout.cuh +++ b/include/tilefoundry/runtime/cuda/layout/shard_layout.cuh @@ -151,6 +151,23 @@ CUTE_HOST_DEVICE constexpr int shard_inner() { cute::tuple_size>::value>{}); } +/// The part of a layout that carries strides. +/// +/// A swizzled layout is a ``cute::ComposedLayout`` whose first component is a +/// function rather than a stride rule, so CuTe deletes ``stride()`` on it +/// (layout_composed.hpp). Every reader of a step reads it off the layout +/// underneath; the function is put back when the tensor is projected. +template +CUTE_HOST_DEVICE constexpr L const &affine_portion(L const &layout) { + return layout; +} + +template +CUTE_HOST_DEVICE constexpr auto +affine_portion(cute::ComposedLayout const &layout) { + return layout.layout_b(); +} + /// Local extent of tensor axis I. template CUTE_HOST_DEVICE constexpr auto local_extent(ShardLayout const &sl) { @@ -175,7 +192,7 @@ CUTE_HOST_DEVICE constexpr auto stride(ShardLayout const &sl) { constexpr int inner = detail::shard_inner, Ax, k>(); return detail::local_extent(sl) * cute::Int{} * - cute::stride(sl.layout_value); + cute::stride(detail::affine_portion(sl.layout_value)); } else { static_assert(detail::attr_leaves_tensor_whole(), "shard layout: this attr must leave the tensor whole"); @@ -220,9 +237,9 @@ CUTE_HOST_DEVICE constexpr auto local_layout(ShardLayout const &sl) { cute::tuple_size::layout{}))>>::value; return [&](std::index_sequence) { - return cute::make_layout( - cute::make_shape(local_extent(sl)...), - cute::make_stride(cute::stride(sl.layout_value)...)); + return cute::make_layout(cute::make_shape(local_extent(sl)...), + cute::make_stride(cute::stride( + affine_portion(sl.layout_value))...)); }(std::make_index_sequence{}); } diff --git a/include/tilefoundry/runtime/cuda/tensor_view/shard_tensor.cuh b/include/tilefoundry/runtime/cuda/tensor_view/shard_tensor.cuh index 9a2b1c24..72ce6041 100644 --- a/include/tilefoundry/runtime/cuda/tensor_view/shard_tensor.cuh +++ b/include/tilefoundry/runtime/cuda/tensor_view/shard_tensor.cuh @@ -41,6 +41,19 @@ template concept ShardTensorLike = detail::is_shard_tensor>::value; +/// Whether a shard layout states a non-affine mapping in front of its layout. +/// +/// A swizzle is one: ``cute::ComposedLayout, Offset, Layout>`` +/// permutes the index the layout produces. ``local_tensor`` keeps that +/// function and folds the instance's origin into the composition's offset, +/// because the function runs on the whole tensor's index: advancing the +/// engine pointer would run it on a slice-relative one and hand two +/// instances one address. The whole-tensor shortcut is wrong for the same +/// reason -- it returns the backing engine, whose layout states no function. +template +inline constexpr bool shard_layout_is_composed_v = + cute::is_composed_layout::value; + /// ``t`` as the tensor this instance holds, in CuTe's ``local_tile`` / /// ``local_partition`` sense: a ShardTensor projected to its own slice, and /// anything the mesh never spread already whole. An op takes both -- an @@ -55,7 +68,9 @@ template CUTE_HOST_DEVICE decltype(auto) local_tensor(T &&t) { if constexpr (!detail::is_shard_tensor::value) { return std::forward(t); } else if constexpr (detail::shard_layout_is_full_broadcast< - typename t_t::shard_layout_type>()) { + typename t_t::shard_layout_type>() && + !shard_layout_is_composed_v< + typename t_t::shard_layout_type>) { return t.engine; } else { auto const [loc_layout, off] = detail::local_layout_and_offset( @@ -65,7 +80,16 @@ template CUTE_HOST_DEVICE decltype(auto) local_tensor(T &&t) { auto &engine_mut = const_cast::type>::type &>( t.engine); - return cute::make_tensor(engine_mut.data() + off, loc_layout); + if constexpr (shard_layout_is_composed_v< + typename t_t::shard_layout_type>) { + auto const &whole = t.shard_layout.layout_value; + return cute::make_tensor( + engine_mut.data(), + cute::make_composed_layout(whole.layout_a(), + whole.offset() + off, loc_layout)); + } else { + return cute::make_tensor(engine_mut.data() + off, loc_layout); + } } } diff --git a/src/tilefoundry/__init__.py b/src/tilefoundry/__init__.py index 192fc16c..285bb204 100644 --- a/src/tilefoundry/__init__.py +++ b/src/tilefoundry/__init__.py @@ -41,6 +41,7 @@ B, Broadcast, ComposedLayout, + Swizzle, Dynamic, IntTuple, Layout, @@ -107,7 +108,7 @@ def view(root, *, port: int = 0, open_browser: bool = True) -> int: "DType", "TensorType", "TupleType", "Type", "Pattern", "DimVarRangePat", "DimVar", - "IntTuple", "LayoutBase", "Layout", "ComposedLayout", + "IntTuple", "LayoutBase", "Layout", "Swizzle", "ComposedLayout", "Topology", "Mesh", "ShardAttr", "Split", "Partial", "Broadcast", "Dynamic", "ShardLayout", "S", "P", "B", diff --git a/src/tilefoundry/codegen/cuda/tir/memory/tensor_view.py b/src/tilefoundry/codegen/cuda/tir/memory/tensor_view.py index 8d24105c..6f01a17a 100644 --- a/src/tilefoundry/codegen/cuda/tir/memory/tensor_view.py +++ b/src/tilefoundry/codegen/cuda/tir/memory/tensor_view.py @@ -25,8 +25,8 @@ from tilefoundry.ir.tir.sync import participation from tilefoundry.ir.types.dim import DimAdd, DimMul, DimSub, DimVar from tilefoundry.ir.types.shape_helpers import shape_numel_upper_bound, upper_bound -from tilefoundry.ir.types.shard import c_order_strides -from tilefoundry.ir.types.shard.layout import ComposedLayout, Layout +from tilefoundry.ir.types.shard import c_order_strides, swizzle_of +from tilefoundry.ir.types.shard.layout import ComposedLayout, Layout, LayoutBase from tilefoundry.ir.types.shard.shard_layout import ( Broadcast, Dynamic, @@ -41,11 +41,58 @@ from tilefoundry.visitor_registry.registries import Role, register_codegen -def _render_layout(shape, strides) -> str: - """Render a CuTe layout type.""" - shape_args = ", ".join(f"cute::Int<{s}>" for s in shape) - stride_args = ", ".join(f"cute::Int<{s}>" for s in strides) - return f"cute::Layout, cute::Stride<{stride_args}>>" +def _swizzle_args(swizzle) -> str: + """``B, M, S`` as the ``cute::Swizzle`` template arguments.""" + return f"{swizzle.bits}, {swizzle.base}, {swizzle.shift}" + + +def _render_layout_type(layout: LayoutBase) -> str: + """Render a CuTe layout type. + + A swizzled composed layout keeps its shape: CuTe states it as a + ``ComposedLayout`` of the swizzle, a static offset and the layout + underneath. The XOR mapping is not expressible as strides, so there is no + affine form to fall back to. + """ + swizzle = swizzle_of(layout) + if swizzle is not None: + return ( + f"cute::ComposedLayout, " + f"cute::Int<{int(layout.offset)}>, {_render_layout_type(layout.outer)}>" + ) + if isinstance(layout, Layout): + shape_args = ", ".join(f"cute::Int<{s}>" for s in layout.shape) + stride_args = ", ".join(f"cute::Int<{s}>" for s in layout.strides) + return f"cute::Layout, cute::Stride<{stride_args}>>" + raise NotImplementedError( + f"tensor_view: no CuTe layout type for {type(layout).__name__}" + ) + + +def _render_layout_value(layout: LayoutBase, dim, stride) -> str: + """Render a CuTe layout value, the mirror of :func:`_render_layout_type`. + + *dim* and *stride* render one shape entry and one stride entry, which is + where a runtime-provided extent reaches the emitted layout. + """ + swizzle = swizzle_of(layout) + if swizzle is not None: + return ( + f"cute::make_composed_layout(" + f"cute::Swizzle<{_swizzle_args(swizzle)}>{{}}, " + f"cute::Int<{int(layout.offset)}>{{}}, " + f"{_render_layout_value(layout.outer, dim, stride)})" + ) + if isinstance(layout, Layout): + shape_args = ", ".join(dim(d) for d in layout.shape) + stride_args = ", ".join(stride(s) for s in layout.strides) + return ( + f"cute::make_layout(cute::make_shape({shape_args}), " + f"cute::make_stride({stride_args}))" + ) + raise NotImplementedError( + f"tensor_view: no CuTe layout value for {type(layout).__name__}" + ) def _scope_mesh_value(mesh, ctx) -> "str | None": @@ -108,7 +155,7 @@ def _render_attr(a) -> str: def _render_shard_layout_type(sl: SL, ctx=None) -> str: """Render a full ShardLayout C++ type string.""" - layout_str = _render_layout(sl.layout.shape, sl.layout.strides) + layout_str = _render_layout_type(sl.layout) attrs_str = ", ".join(_render_attr(a) for a in sl.attrs) mesh_str = _render_mesh_type(sl.mesh, ctx) return f"tilefoundry::ShardLayout<{layout_str}, cute::tuple<{attrs_str}>, {mesh_str}>" @@ -163,6 +210,12 @@ def render_shard_layout_value( """ sll = sl.layout if storage is StorageKind.RMEM: + if swizzle_of(sll) is not None: + raise NotImplementedError( + "render_shard_layout_value: a Swizzle states how a shared-memory " + "bank pattern is arranged; a register engine has no such addresses, " + "and rebuilding this layout from register strides would drop it" + ) sll = Layout(shape=sll.shape, strides=register_strides(sl)) mesh_layout = sl.mesh.layout if isinstance(mesh_layout, ComposedLayout): @@ -234,8 +287,6 @@ def _mesh_dim(d): ml_var = f"{var_name}__mesh_layout" mesh_var = f"{var_name}__mesh" - sl_shape_args = ", ".join(_global_dim(d) for d in sll.shape) - sl_stride_args = ", ".join(_static_dim(s, "shard layout stride") for s in sll.strides) ml_shape_args = ", ".join(_mesh_dim(d) for d in ml_shape) ml_stride_args = ", ".join(_static_dim(s, "mesh layout stride") for s in ml_strides) @@ -246,10 +297,10 @@ def _mesh_dim(d): mesh_layout = _composed_mesh_layout(positions, ml_base) attrs = ", ".join(_render_attr(a) for a in sl.attrs) - preamble = [ - f"auto {sl_var} = cute::make_layout(" - f"cute::make_shape({sl_shape_args}), cute::make_stride({sl_stride_args}));", - ] + sl_layout = _render_layout_value( + sll, _global_dim, lambda s: _static_dim(s, "shard layout stride") + ) + preamble = [f"auto {sl_var} = {sl_layout};"] scope_mesh = _scope_mesh_value(sl.mesh, ctx) if scope_mesh is not None: mesh_var = scope_mesh diff --git a/src/tilefoundry/inspection/printer_base.py b/src/tilefoundry/inspection/printer_base.py index 149c0cbc..57535686 100644 --- a/src/tilefoundry/inspection/printer_base.py +++ b/src/tilefoundry/inspection/printer_base.py @@ -20,7 +20,7 @@ DimSub, DimVar, ) -from tilefoundry.ir.types.shard.layout import ComposedLayout, Layout, LayoutBase +from tilefoundry.ir.types.shard.layout import ComposedLayout, Layout, LayoutBase, Swizzle from tilefoundry.ir.types.shard.mesh import Mesh from tilefoundry.ir.types.shard.shard_layout import Broadcast, Partial, ShardLayout, Split from tilefoundry.ir.types.storage import StorageKind @@ -290,6 +290,11 @@ def visit_Layout(self, value: Layout, ctx=None) -> str: strides = self.shape_tuple(value.strides, ctx) if value.strides is not None else "None" return f"Layout({self.shape_tuple(value.shape, ctx)}, {strides})" + def visit_Swizzle(self, value: Swizzle, ctx=None) -> str: + if ctx is not None: + ctx.use(PythonExpr(("from tilefoundry.ir.types.shard import Swizzle",), "")) + return f"Swizzle({value.bits}, {value.base}, {value.shift})" + def visit_ComposedLayout(self, value: ComposedLayout, ctx=None) -> str: if ctx is not None: ctx.use(PythonExpr(("from tilefoundry.ir.types.shard import ComposedLayout",), "")) diff --git a/src/tilefoundry/ir/hir/tensor/slice.py b/src/tilefoundry/ir/hir/tensor/slice.py index 5bbd5f2e..8d6ecddb 100644 --- a/src/tilefoundry/ir/hir/tensor/slice.py +++ b/src/tilefoundry/ir/hir/tensor/slice.py @@ -24,6 +24,7 @@ ComposedLayout, Layout, ShardLayout, + Swizzle, ) from tilefoundry.ir.types.shard.shard_layout import ( layout_axis_to_tensor_axis, @@ -378,6 +379,16 @@ def _slice_shard_layout(call, ctx, x_ty, source, starts, inherited_offset): @register_typeinfer(Slice) def _(call: "Call", ctx: "TypeInferContext") -> TensorType: + """The window's type, its source layout carried where the window keeps it. + + A ``Swizzle`` survives a window. It says where an element lives, not which + element a coordinate names, so it leaves the domain -- and the access + relation stated over it -- alone; the window narrows that domain and + shifts the index by a constant the composition's ``offset`` already holds. + Any other composed ``inner`` is refused by name: a window cannot be proven + against a mapping this does not know, and reporting the outer strides as + the whole layout would be wrong rather than partial. + """ x_ty = ctx.type_of(call.args[0]) starts = call.args[1] op = call.target @@ -400,13 +411,29 @@ def _(call: "Call", ctx: "TypeInferContext") -> TensorType: layout_shape = shape source = x_ty.layout inherited_offset = 0 - if ( - isinstance(source, ComposedLayout) - and source.inner is None - and isinstance(source.outer, (Layout, ShardLayout)) + inherited_inner = None + if isinstance(source, ComposedLayout) and isinstance( + source.outer, (Layout, ShardLayout) ): + if isinstance(source.inner, Swizzle) and isinstance(source.outer, Layout): + inherited_inner = source.inner + elif source.inner is not None: + ctx.error( + call, + f"Slice cannot narrow a composed layout whose inner is " + f"{type(source.inner).__name__} over " + f"{type(source.outer).__name__}: the window would have to be " + f"proven against that mapping, not against the outer strides", + ) inherited_offset = source.offset source = source.outer + elif isinstance(source, ComposedLayout) and source.inner is not None: + ctx.error( + call, + f"Slice cannot narrow a composed layout whose inner is " + f"{type(source.inner).__name__}: the window would have to be proven " + f"against that mapping, not against the outer strides", + ) new_layout = None if isinstance(source, ShardLayout): @@ -428,7 +455,7 @@ def _(call: "Call", ctx: "TypeInferContext") -> TensorType: steps.append(stride) else: new_layout = ComposedLayout( - inner=None, + inner=inherited_inner, offset=inherited_offset + sum( start * stride diff --git a/src/tilefoundry/ir/types/shard/__init__.py b/src/tilefoundry/ir/types/shard/__init__.py index 06b871ac..9ed5d14a 100644 --- a/src/tilefoundry/ir/types/shard/__init__.py +++ b/src/tilefoundry/ir/types/shard/__init__.py @@ -3,8 +3,14 @@ # ruff: noqa: I001 -- curated re-export order; alphabetical sort breaks staged imports. from .int_tuple import IntTuple, flatten, product -from .layout import ComposedLayout, Layout, LayoutBase -from .layout_algebra import c_order_strides, prefix_product, try_c_order_strides +from .layout import ComposedLayout, Layout, LayoutBase, Swizzle +from .layout_algebra import ( + c_order_strides, + composition, + prefix_product, + swizzle_of, + try_c_order_strides, +) from .mesh import ( Mesh, Topology, @@ -40,7 +46,10 @@ "prefix_product", "LayoutBase", "Layout", + "Swizzle", "ComposedLayout", + "composition", + "swizzle_of", "Topology", "check_topology", "composed", diff --git a/src/tilefoundry/ir/types/shard/layout.py b/src/tilefoundry/ir/types/shard/layout.py index 54963902..665ba70f 100644 --- a/src/tilefoundry/ir/types/shard/layout.py +++ b/src/tilefoundry/ir/types/shard/layout.py @@ -22,25 +22,88 @@ class Layout(LayoutBase): strides: Optional[tuple["ShapeDim", ...]] = None +@dataclass(frozen=True) +class Swizzle: + """CuTe ``Swizzle``: the XOR offset functor, not a layout. + + It states a permutation of an *index*, so it has no domain shape of its + own and is not a ``LayoutBase``. The domain comes from the + ``ComposedLayout`` that carries it in ``inner``. + + ``bits`` is CuTe's ``BBits`` (how many bits are XORed), ``base`` its + ``MBase`` (how many low bits are left alone) and ``shift`` its ``SShift`` + (the signed distance from the source mask to the target mask). + + See [shard §4.1](docs/spec/shard.md#41-swizzle). + """ + + bits: int + base: int + shift: int + + def __post_init__(self) -> None: + for name in ("bits", "base", "shift"): + value = getattr(self, name) + if not isinstance(value, int) or isinstance(value, bool): + raise ValueError(f"Swizzle {name} must be an int, got {value!r}") + if self.bits < 0: + raise ValueError(f"Swizzle bits must be non-negative, got {self.bits}") + if self.base < 0: + raise ValueError(f"Swizzle base must be non-negative, got {self.base}") + if abs(self.shift) < self.bits: + raise ValueError( + f"abs(Swizzle shift) must be >= bits, got shift={self.shift}, " + f"bits={self.bits}" + ) + + @property + def bit_mask(self) -> int: + """CuTe ``bit_msk``: the ``bits`` low bits.""" + return (1 << self.bits) - 1 + + @property + def yyy_mask(self) -> int: + """CuTe ``yyy_msk``: the bits read out of the index.""" + return self.bit_mask << (self.base + max(0, self.shift)) + + @property + def zzz_mask(self) -> int: + """CuTe ``zzz_msk``: the bits XORed in the index.""" + return self.bit_mask << (self.base - min(0, self.shift)) + + @property + def swizzle_code(self) -> int: + """CuTe ``swizzle_code``: every bit this swizzle touches.""" + return self.yyy_mask | self.zzz_mask + + def __call__(self, offset: int) -> int: + """``offset ^ shiftr(offset & yyy_msk, msk_sft)`` (CuTe ``apply``).""" + selected = offset & self.yyy_mask + moved = selected >> self.shift if self.shift >= 0 else selected << -self.shift + return offset ^ moved + + @dataclass(frozen=True) class ComposedLayout(LayoutBase): """Represent ``image(c) = inner(offset + outer(c))``. ``outer`` defines the domain shape and axis numbering; ``None`` means an - identity component. Either component may retain a nested ``ShardLayout``. - Inversion applies the component inverses in reverse order. + identity component. ``outer`` and a ``LayoutBase`` ``inner`` may retain a + nested ``ShardLayout``. Inversion applies the component inverses in reverse + order. ``inner`` may also be a :class:`Swizzle`, which states a + non-affine permutation of the index and carries no domain of its own. See [shard §4](docs/spec/shard.md#4-composedlayout). """ - inner: LayoutBase | None + inner: "LayoutBase | Swizzle | None" offset: int outer: LayoutBase | None @property def shape(self) -> tuple: domain = self.outer if self.outer is not None else self.inner - if domain is None: + if domain is None or isinstance(domain, Swizzle): return () return domain.shape @@ -48,4 +111,4 @@ def shape(self) -> tuple: EMPTY_LAYOUT = Layout(shape=(), strides=()) -__all__ = ["LayoutBase", "Layout", "ComposedLayout", "EMPTY_LAYOUT"] +__all__ = ["LayoutBase", "Layout", "Swizzle", "ComposedLayout", "EMPTY_LAYOUT"] diff --git a/src/tilefoundry/ir/types/shard/layout_algebra.py b/src/tilefoundry/ir/types/shard/layout_algebra.py index 4447c675..43e7985a 100644 --- a/src/tilefoundry/ir/types/shard/layout_algebra.py +++ b/src/tilefoundry/ir/types/shard/layout_algebra.py @@ -1,8 +1,9 @@ """Provide flat CuTe layout algebra for mesh execution scopes. The restricted port supports coordinate application, inverses, containment, -and projection for ``Layout`` and ``ComposedLayout``. Execution scopes must be -injective and inverse-projectable. +and projection for ``Layout`` and ``ComposedLayout``, and the CuTe +``swizzle_layout.hpp`` specializations for a ``Swizzle`` in ``inner``. +Execution scopes must be injective and inverse-projectable. See [shard §9](docs/spec/shard.md#9-layout-construction-and-mesh-scope-projection). """ @@ -12,7 +13,7 @@ from typing import Optional, Union from .int_tuple import flatten, product -from .layout import ComposedLayout, Layout +from .layout import ComposedLayout, Layout, Swizzle class NotProjectable(ValueError): @@ -74,8 +75,26 @@ def size(layout: Layout) -> int: return product(layout.shape) -def apply(layout: Layout, coord: int) -> int: - """``crd2idx`` of a 1-D domain coord: decompose by shape, dot with strides.""" +def swizzle_of(layout: object) -> Optional[Swizzle]: + """The ``Swizzle`` a composed layout applies last, or ``None``. + + CuTe ``get_swizzle_portion``, answering ``None`` rather than the identity + ``Swizzle<0,4,3>`` so a caller can branch on "is this swizzled at all". + """ + if isinstance(layout, ComposedLayout) and isinstance(layout.inner, Swizzle): + return layout.inner + return None + + +def apply(layout: Union[Layout, ComposedLayout], coord: int) -> int: + """``crd2idx`` of a 1-D domain coord: decompose by shape, dot with strides. + + A ``ComposedLayout`` applies its components in order, so this is + ``inner(offset + outer(coord))`` — the swizzle case included, since a + ``Swizzle`` is exactly a mapping on that index. + """ + if isinstance(layout, ComposedLayout): + return _apply_any(layout, coord) shape = _shape(layout) stride = _stride(layout) idx = 0 @@ -86,7 +105,38 @@ def apply(layout: Layout, coord: int) -> int: return idx -def cosize(layout: Layout) -> int: +def _crd2idx_unbounded(layout: Layout, coord: int) -> int: + """``apply`` as CuTe's ``crd2idx`` has it: the last mode is not wrapped. + + ``apply`` wraps every mode, which is the same answer for an in-domain + coord and the one this module's callers want. The swizzle composition + rules feed a *bit mask* through a layout instead, which is routinely + larger than the domain, and CuTe leaves the final mode unwrapped so those + high bits keep contributing. Only that port reads this. + """ + shape = _shape(layout) + stride = _stride(layout) + idx = 0 + rem = coord + last = len(shape) - 1 + for position, (s, d) in enumerate(zip(shape, stride)): + if position == last: + idx += rem * d + else: + idx += (rem % s) * d + rem //= s + return idx + + +def cosize(layout: Union[Layout, ComposedLayout]) -> int: + """The codomain extent. + + A swizzle permutes bits inside the codomain its ``outer`` already spans, + so it adds nothing to it: CuTe ``cosize`` of a swizzled composed layout is + ``cosize`` of the layout underneath (``swizzle_layout.hpp:172``). + """ + if swizzle_of(layout) is not None: + return cosize(layout.outer) return apply(layout, size(layout) - 1) + 1 @@ -123,8 +173,16 @@ def unflatten(flat: tuple, profile) -> tuple: return nested -def coalesce(layout: Layout) -> Layout: - """Flatten + merge contiguous modes, drop shape-1 modes (CuTe ``coalesce``).""" +def coalesce(layout: Union[Layout, ComposedLayout]): + """Flatten + merge contiguous modes, drop shape-1 modes (CuTe ``coalesce``). + + Coalescing renames the domain and leaves the index mapping alone, so a + swizzled composed layout coalesces underneath its swizzle. + """ + if swizzle_of(layout) is not None: + return ComposedLayout( + inner=layout.inner, offset=layout.offset, outer=coalesce(layout.outer) + ) result_shape: list[int] = [1] result_stride: list[int] = [0] for shape, stride in zip(_shape(layout), _stride(layout)): @@ -239,7 +297,88 @@ def _check_admissible(scope: ComposedLayout) -> None: raise NotProjectable("outer layout is not inverse-projectable (injective + compact)") -def left_inverse(layout: Union[Layout, ComposedLayout]): +def _make_swizzle(active_y: int, active_z: int) -> Swizzle: + """CuTe ``make_swizzle()``: the swizzle that XORs *Y* onto *Z*. + + The two masks must hold the same number of bits; their trailing-zero + counts give ``base`` and the signed ``shift``, and the reconstructed + ``swizzle_code`` must give the masks back, which is how CuTe checks that + the pair is a swizzle it can represent at all. + """ + bits_y, bits_z = active_y.bit_count(), active_z.bit_count() + if bits_y != bits_z: + raise NotImplementedError( + f"composition: the Y mask {active_y:#x} holds {bits_y} bits and the Z " + f"mask {active_z:#x} holds {bits_z}; only an equal-width pair is a " + f"Swizzle" + ) + if bits_y == 0: + return Swizzle(0, 0, 0) + trailing_y = (active_y & -active_y).bit_length() - 1 + trailing_z = (active_z & -active_z).bit_length() - 1 + swizzle = Swizzle(bits_y, min(trailing_y, trailing_z), trailing_y - trailing_z) + if swizzle.swizzle_code != (active_y | active_z): + raise NotImplementedError( + f"composition: the mask pair ({active_y:#x}, {active_z:#x}) is not a " + f"Swizzle; its bits are not two contiguous equal-width runs" + ) + return swizzle + + +def composition(left, right, offset: int = 0): + """CuTe ``composition`` for the swizzle cases (``swizzle_layout.hpp:302``). + + ``composition(Swizzle, Layout)`` builds a swizzled layout, which states + ``Swizzle(offset + Layout(coord))``. + + ``composition(Layout, Swizzle)`` would otherwise want the ``Swizzle`` in + ``outer``, which has no domain to be a domain-side component of. CuTe + instead reads which of the swizzle's bits the layout leaves active, + rebuilds a ``Swizzle`` over those, and puts it back on the inner side. + """ + if isinstance(left, Swizzle) and isinstance(right, Layout): + if left.bits == 0 and offset == 0: + return right + return ComposedLayout(inner=left, offset=offset, outer=right) + if isinstance(left, Layout) and isinstance(right, Swizzle): + if offset: + raise NotImplementedError( + f"composition: a non-zero offset ({offset}) between a Layout and a " + f"Swizzle has no canonical ComposedLayout form" + ) + active_y = _crd2idx_unbounded(left, right.yyy_mask) + active_z = _crd2idx_unbounded(left, right.zzz_mask) + return composition(_make_swizzle(active_y, active_z), left) + raise NotImplementedError( + f"composition: no rule for {type(left).__name__} ∘ {type(right).__name__}" + ) + + +def _swizzled_inverse(layout: ComposedLayout, inverse_of_layout): + """CuTe's swizzled ``left_inverse``/``right_inverse`` (``swizzle_layout.hpp:344``). + + ``inverse(Swizzle(offset + outer(c)))`` passes the swizzle back to the + left of the inverted ``outer``, which ``composition(Layout, Swizzle)`` + then canonicalizes back into this IR's one legal shape. CuTe's non-zero + ``offset`` branch composes ``inverse(offset)`` between the two, which + lands a bare ``Swizzle`` in ``outer``; that is not a layout, so this + refuses it by name rather than building something unrepresentable. + """ + if layout.offset != 0: + raise NotImplementedError( + f"inverse: a swizzled composed layout with a non-zero offset " + f"({layout.offset}) inverts to a Swizzle on the domain side, which " + f"ComposedLayout.outer cannot hold" + ) + if not isinstance(layout.outer, Layout): + raise NotImplementedError( + f"inverse: a swizzled composed layout inverts through its outer " + f"Layout; this one states {type(layout.outer).__name__}" + ) + return composition(inverse_of_layout(layout.outer), layout.inner) + + +def left_inverse(layout: Union[Layout, ComposedLayout, Swizzle]): """CuTe ``left_inverse``, dispatched. Plain ``Layout`` → the flat algebra. ``ComposedLayout`` → the recursive @@ -249,6 +388,10 @@ def left_inverse(layout: Union[Layout, ComposedLayout]): outer=None)`` (``outer=None`` ≡ identity), i.e. ``image⁻¹(t) = outer⁻¹(t − offset)``. """ + if isinstance(layout, Swizzle): + return layout + if swizzle_of(layout) is not None: + return _swizzled_inverse(layout, left_inverse) if isinstance(layout, ComposedLayout): _check_admissible(layout) return ComposedLayout( @@ -259,8 +402,16 @@ def left_inverse(layout: Union[Layout, ComposedLayout]): return _left_inverse_layout(layout) -def right_inverse(layout: Union[Layout, ComposedLayout]): - """CuTe ``right_inverse``, dispatched (mirror of :func:`left_inverse`).""" +def right_inverse(layout: Union[Layout, ComposedLayout, Swizzle]): + """CuTe ``right_inverse``, dispatched (mirror of :func:`left_inverse`). + + A ``Swizzle`` is an involution -- its Y and Z bit ranges do not overlap -- + so it is its own inverse on both sides (``swizzle_layout.hpp:371``). + """ + if isinstance(layout, Swizzle): + return layout + if swizzle_of(layout) is not None: + return _swizzled_inverse(layout, right_inverse) if isinstance(layout, ComposedLayout): _check_admissible(layout) return ComposedLayout( @@ -275,6 +426,8 @@ def _apply_any(layout, x: int) -> int: """Apply a ``Layout`` / ``ComposedLayout`` (``None`` ≡ identity) to ``x``.""" if layout is None: return x + if isinstance(layout, Swizzle): + return layout(x) if isinstance(layout, Layout): return apply(layout, x) if isinstance(layout, ComposedLayout): @@ -324,6 +477,8 @@ def contains(scope: ComposedLayout, t: int) -> bool: "NotProjectable", "prefix_product", "size", + "swizzle_of", + "composition", "cosize", "apply", "idx2crd", diff --git a/tests/ops/tir/cuda/test_swizzle.py b/tests/ops/tir/cuda/test_swizzle.py new file mode 100644 index 00000000..ca42d239 --- /dev/null +++ b/tests/ops/tir/cuda/test_swizzle.py @@ -0,0 +1,124 @@ +"""Carry a CuTe XOR swizzle from the IR to shared memory and back. + +See [shard §4.1](docs/spec/shard.md#41-swizzle). +""" + +from __future__ import annotations + +import pytest +import torch + +import tilefoundry +import tilefoundry.codegen.cuda # noqa: F401 -- trigger emitter autodiscovery +from tests._source import import_dsl +from tilefoundry import module, prim_func +from tilefoundry.codegen.cuda.tir.memory.tensor_view import render_shard_layout_value +from tilefoundry.dsl import T, Tensor +from tilefoundry.inspection import as_script +from tilefoundry.ir.core.kinds import BinaryKind +from tilefoundry.ir.types.shard import ( + ComposedLayout, + Layout, + Mesh, + ShardLayout, + Split, + Swizzle, + Topology, +) +from tilefoundry.target import CpuTarget, CudaTarget + +_CUDA = CudaTarget("nvidia.h200_sxm") +_ROWS, _COLS = 128, 4 +_SWIZZLE = Swizzle(2, 2, 2) +_THREADS = Mesh((Topology("thread", _ROWS),), Layout((_ROWS,), (1,)), ("t",)) + + +def _rows(mesh: Mesh) -> ShardLayout: + """One row per thread, addressed as the tile is laid out.""" + return ShardLayout(Layout((_ROWS, _COLS), (_COLS, 1)), (Split(0),), mesh) + + +def _swizzled_rows(mesh: Mesh) -> ShardLayout: + """The same rows, reached through the swizzle. + + ``Swizzle(2, 2, 2)`` XORs Y bits 4-5 onto Z bits 2-3, so over this + (128, 4) f32 tile thread ``t`` holds the row ``t ^ ((t >> 2) & 3)``: the + address really moves, and the four elements of a row stay together. + """ + return ShardLayout( + ComposedLayout( + inner=_SWIZZLE, offset=0, outer=Layout((_ROWS, _COLS), (_COLS, 1)) + ), + (Split(0),), + mesh, + ) + + +@module(entry="swizzled_square_host", target=_CUDA, topologies=(Topology("thread", _ROWS),)) +class SwizzledSquare: + """Square a tile that is staged through swizzled shared memory.""" + + @prim_func(target=_CUDA) + def swizzled_square_device( + src: Tensor[(_ROWS, _COLS), "f32"], + dst: Tensor[(_ROWS, _COLS), "f32"], + ): + with Mesh((Topology("thread", _ROWS),), Layout((_ROWS,), (1,)), ("t",)) as threads: + src_view = T.tensor_view(src, layout=_rows(threads)) + dst_view = T.tensor_view(dst, layout=_rows(threads)) + tile = T.alloc_tensor( + Tensor[(_ROWS, _COLS), "f32", _swizzled_rows(threads), "smem"] + ) + fragment = T.alloc_tensor( + Tensor[(_ROWS, _COLS), "f32", _rows(threads), "rmem"] + ) + T.copy(src_view, tile) + T.sync(threads) + T.copy(tile, fragment) + T.binary(fragment, fragment, fragment, kind=BinaryKind.MUL) + T.copy(fragment, dst_view) + + + @prim_func(target=CpuTarget()) + def swizzled_square_host( + src: Tensor[(_ROWS, _COLS), "f32"], + dst: Tensor[(_ROWS, _COLS), "f32"], + ): + launch( # noqa: F821 + swizzled_square_device, # noqa: F821 + src, + dst, + grid=(1, 1, 1), + block=(_ROWS, 1, 1), + ) + + +@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires CUDA") +def test_swizzled_smem_load_compiles_and_executes() -> None: + """One kernel through print, parse, codegen, nvcc and the device. + + The emitted layout is asserted as well, because running the kernel cannot + tell us about it: a swizzle permutes addresses, so writing and reading + through the same layout returns the same numbers whether it is applied or + dropped entirely. + """ + source = as_script(SwizzledSquare) + restored = import_dsl(source, name="SwizzledSquare") + assert as_script(restored) == source + + preamble, _ = render_shard_layout_value("tile", _swizzled_rows(_THREADS)) + assert any( + "cute::make_composed_layout(cute::Swizzle<2, 2, 2>{}, cute::Int<0>{}, " + "cute::make_layout(cute::make_shape(cute::Int<128>{}, cute::Int<4>{}), " + "cute::make_stride(cute::Int<4>{}, cute::Int<1>{})))" in line + for line in preamble + ), preamble + + runtime_module = tilefoundry.compile(restored, target=_CUDA) + torch.manual_seed(0) + src = torch.randn(_ROWS, _COLS, dtype=torch.float32, device="cuda") + dst = torch.zeros_like(src) + runtime_module(src, dst) + torch.cuda.synchronize() + + assert torch.equal(dst, src * src)