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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions docs/spec/runtime.md
Original file line number Diff line number Diff line change
Expand Up @@ -1044,6 +1044,9 @@ template <class T>
concept ShardTensorLike =
detail::is_shard_tensor<cute::remove_cvref_t<T>>::value;

template <class SL>
inline constexpr bool shard_layout_is_composed_v;

template <class T> CUTE_HOST_DEVICE decltype(auto) local_tensor(T &&t);

template <class T>
Expand All @@ -1058,6 +1061,8 @@ template <class T> 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.
Expand Down
64 changes: 61 additions & 3 deletions docs/spec/shard.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
```
Expand All @@ -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:

Expand All @@ -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<B,M,S>`: the XOR offset functor a swizzled
shared-memory layout applies to its index.

```python
class Swizzle:
"""Describe the XOR permutation CuTe `Swizzle<B,M,S>` 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`
Expand Down
25 changes: 21 additions & 4 deletions include/tilefoundry/runtime/cuda/layout/shard_layout.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -151,6 +151,23 @@ CUTE_HOST_DEVICE constexpr int shard_inner() {
cute::tuple_size<shard_mesh_flat_t<SL>>::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 <class L>
CUTE_HOST_DEVICE constexpr L const &affine_portion(L const &layout) {
return layout;
}

template <class A, class O, class B>
CUTE_HOST_DEVICE constexpr auto
affine_portion(cute::ComposedLayout<A, O, B> const &layout) {
return layout.layout_b();
}

/// Local extent of tensor axis I.
template <size_t I, class L, class A, class M>
CUTE_HOST_DEVICE constexpr auto local_extent(ShardLayout<L, A, M> const &sl) {
Expand All @@ -175,7 +192,7 @@ CUTE_HOST_DEVICE constexpr auto stride(ShardLayout<L, A, M> const &sl) {
constexpr int inner =
detail::shard_inner<ShardLayout<L, A, M>, Ax, k>();
return detail::local_extent<size_t(k)>(sl) * cute::Int<inner>{} *
cute::stride<k>(sl.layout_value);
cute::stride<k>(detail::affine_portion(sl.layout_value));
} else {
static_assert(detail::attr_leaves_tensor_whole<attr_t>(),
"shard layout: this attr must leave the tensor whole");
Expand Down Expand Up @@ -220,9 +237,9 @@ CUTE_HOST_DEVICE constexpr auto local_layout(ShardLayout<L, A, M> const &sl) {
cute::tuple_size<cute::remove_cvref_t<decltype(cute::shape(
typename ShardLayout<L, A, M>::layout{}))>>::value;
return [&]<size_t... Is>(std::index_sequence<Is...>) {
return cute::make_layout(
cute::make_shape(local_extent<Is>(sl)...),
cute::make_stride(cute::stride<int(Is)>(sl.layout_value)...));
return cute::make_layout(cute::make_shape(local_extent<Is>(sl)...),
cute::make_stride(cute::stride<int(Is)>(
affine_portion(sl.layout_value))...));
}(std::make_index_sequence<t_rank>{});
}

Expand Down
28 changes: 26 additions & 2 deletions include/tilefoundry/runtime/cuda/tensor_view/shard_tensor.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,19 @@ template <class T>
concept ShardTensorLike =
detail::is_shard_tensor<cute::remove_cvref_t<T>>::value;

/// Whether a shard layout states a non-affine mapping in front of its layout.
///
/// A swizzle is one: ``cute::ComposedLayout<Swizzle<B,M,S>, 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 <class SL>
inline constexpr bool shard_layout_is_composed_v =
cute::is_composed_layout<typename SL::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
Expand All @@ -55,7 +68,9 @@ template <class T> CUTE_HOST_DEVICE decltype(auto) local_tensor(T &&t) {
if constexpr (!detail::is_shard_tensor<t_t>::value) {
return std::forward<T>(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(
Expand All @@ -65,7 +80,16 @@ template <class T> CUTE_HOST_DEVICE decltype(auto) local_tensor(T &&t) {
auto &engine_mut = const_cast<typename std::remove_const<
typename std::remove_reference<decltype(t.engine)>::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);
}
}
}

Expand Down
3 changes: 2 additions & 1 deletion src/tilefoundry/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
B,
Broadcast,
ComposedLayout,
Swizzle,
Dynamic,
IntTuple,
Layout,
Expand Down Expand Up @@ -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",
Expand Down
79 changes: 65 additions & 14 deletions src/tilefoundry/codegen/cuda/tir/memory/tensor_view.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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::Shape<{shape_args}>, 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<cute::Swizzle<{_swizzle_args(swizzle)}>, "
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::Shape<{shape_args}>, 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":
Expand Down Expand Up @@ -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}>"
Expand Down Expand Up @@ -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):
Expand Down Expand Up @@ -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)

Expand All @@ -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
Expand Down
7 changes: 6 additions & 1 deletion src/tilefoundry/inspection/printer_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",), ""))
Expand Down
Loading
Loading