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 trait
MHATileScheduler
Describes a schedule for the persistent MHA kernel.
A tile scheduler maps work tiles to thread blocks, advances the per-CTA state through the work grid across kernel iterations, and reports the grid dimensions required for launch.
Implemented traits
AnyType,
Copyable,
Deinitable,
DevicePassable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
device_type
comptime device_type
Indicate the type being used on accelerator devices.
may_advance
comptime may_advance
mha_schedule
comptime mha_schedule
Required methods
__init__
def __init__(out self, *, copy: Self)
Create a new instance of the value by copying an existing one.
Args:
- copy (
_Self): The value to copy.
Returns:
_Self
def __init__(out self, *, deinit move: Self)
Create a new instance of the value by moving the value of another.
Args:
- move (
_Self): The value to move.
Returns:
_Self
get_current_work_info
def get_current_work_info[ValidLengthType: OptionalPointer, //](self, ts: MHATileSummary[ValidLengthType], state: MHATileState) -> WorkInfo
Returns the current WorkInfo.
Parameters:
- ValidLengthType (
OptionalPointer): The optional pointer type carrying per-batch sequence length offsets (inferred).
Args:
- ts (
MHATileSummary[ValidLengthType]): The tile summary describing the work grid. - state (
MHATileState): The per-CTA scheduler state whose current index to resolve.
Returns:
advance
def advance[ValidLengthType: OptionalPointer, //, producer: Bool, sync: MHASchedulerSynchronization = MHASchedulerSynchronization.DEFAULT](self, ts: MHATileSummary[ValidLengthType], mut state: MHATileState, pipeline_idx: UInt32) -> OptionalReg[SeqInfo]
Advance state to the next work item.
func must return a Bool indicating whether there is more work.
Returns True if there is more work.
Parameters:
- ValidLengthType (
OptionalPointer): The optional pointer type carrying per-batch sequence length offsets (inferred). - producer (
Bool): Whether the calling CTA is the producer thread for copy-async paths. - sync (
MHASchedulerSynchronization): Which threads participate in the barrier when advancing (defaults toMHASchedulerSynchronization.DEFAULT).
Args:
- ts (
MHATileSummary[ValidLengthType]): The tile summary describing the work grid. - state (
MHATileState): The mutable per-CTA scheduler state to advance. - pipeline_idx (
UInt32): The pipeline stage index for storing the shared work index.
Returns:
grid_dim
static def grid_dim(batch_size: UInt32, max_num_prompt_tiles: UInt32) -> Tuple[Int, Int, Int]
Return the grid_dim required for the kernel.
Args:
- batch_size (
UInt32): Number of sequences in the batch. - max_num_prompt_tiles (
UInt32): Maximum number of prompt tiles along the sequence dimension.
Returns:
initial_state
def initial_state[ValidLengthType: OptionalPointer, //](self, ptr: Pointer[UInt32, MutAnyOrigin, address_space=AddressSpace.SHARED], tile_summary: MHATileSummary[ValidLengthType]) -> MHATileState
Create the initial state object.
Parameters:
- ValidLengthType (
OptionalPointer): The optional pointer type carrying per-batch sequence length offsets (inferred).
Args:
- ptr (
Pointer[UInt32, MutAnyOrigin, address_space=AddressSpace.SHARED]): Shared-memory pointer for communicating the active work index across threads. - tile_summary (
MHATileSummary[ValidLengthType]): The tile summary describing the work grid dimensions.
Returns:
unsafe_seq_info
def unsafe_seq_info[ValidLengthType: OptionalPointer, //](self, ts: MHATileSummary[ValidLengthType], state: MHATileState) -> SeqInfo
Returns:
get_type_name
static def get_type_name() -> String
Gets the name of the host type (the one implementing this trait). For example, Int would return "Int", DeviceBuffer[DType.float32] would return "DeviceBuffer[DType.float32]". This is used for error messages when passing types to the device. TODO: This method will be retired soon when better kernel call error messages arrive.
Returns:
String: The host type's name.
Provided methods
copy
def copy(self) -> Self
Explicitly construct a copy of self, a convenience method for Self(copy=self) when the type is inconvenient to write out.
Overriding this method is not allowed.
Returns:
_Self: A copy of this value.