IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Mojo struct

EPRoleSplit

struct EPRoleSplit[block_size: Int, n_items: Int, flag: Bool = False]

Splits one dispatch CTA into a copy role and a publisher role.

The copy role carries the whole copy/quantize body for a token. The publisher role issues no payload and no scale stores at all, so its warps reach the send and the completion publication with an empty store history instead of first draining a full token's worth of their own stores. That empty history is the point of the split: it is what a CTA with more warps than the copy loop needs gets for free, and what a CTA sized to the copy loop does not.

Every role test in the dispatch path must come from here. A body that re-derives its role from a bare warp_id() comparison or a literal thread count can drift out of agreement with the other half of the split.

Parameters​

  • ​block_size (Int): Number of threads in the dispatch CTA.
  • ​n_items (Int): Number of copy items in one token.
  • ​flag (Bool): Caller opt-in. Off by default, which forces enabled false and sends every consumer back to the stock strided body. The geometry predicate below is an ADDITIONAL guard, never a substitute: both the flag and the geometry must hold.

Implemented traits​

AnyType, Deinitable, Movable

comptime members​

enabled​

comptime enabled = flag and (Int((add (block_size // _resolve_warp_size()), -4)) >= Int(1)) and (n_items == Int((add (mul (block_size // _resolve_warp_size()), _resolve_warp_size(), 3), (mul _resolve_warp_size(), -12))))

n_copy_threads​

comptime n_copy_threads = (Int((add (block_size // _resolve_warp_size()), -4)) * _resolve_warp_size())

n_copy_warps​

comptime n_copy_warps = ((block_size // _resolve_warp_size()) - Int(4))

n_publisher_warps​

comptime n_publisher_warps = 4

n_trips​

comptime n_trips = 3

n_warps​

comptime n_warps = (block_size // _resolve_warp_size())

Methods​

copy_role_index​

static def copy_role_index() -> Int

Returns this thread's linear index within the copy role.

Returns:

Int

is_ep_copy_role​

static def is_ep_copy_role() -> Bool

Returns True if this thread carries copy/quantize work.

Returns:

Bool

publisher_role_index​

static def publisher_role_index() -> Int

Returns this warp's index within the publisher role.

Negative on a copy-role warp, so it is only meaningful behind is_ep_publisher_role().

Returns:

Int

is_ep_publisher_role​

static def is_ep_publisher_role() -> Bool

Returns True if this thread's warp carries publication work.

Returns:

Bool

Was this page helpful?