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
QueuedTileScheduler
struct QueuedTileScheduler[tile_shape: UInt32, num_heads: UInt32, /, decoding: Bool, num_ctas: UInt32 = SIMD(GPUInfo.from_family(AcceleratorArchitectureFamily(Int(32), Int(2048), Int(233472), Int(65536), Int(1024)), StringSpan("H100"), StringSpan("cuda"), StringSpan("hopper"), SIMD(9), StringSpan("sm_90a"), Int(132)).sm_count), schedule: MHASchedule = MHASchedule.DEFAULT]
If decoding == False, then num_heads is q_num_heads. If decoding == True, then num_heads is kv_num_heads.
Parameters
- tile_shape (
UInt32): Size of each query tile along the sequence dimension. - num_heads (
UInt32): Number of attention heads (q_num_headswhen not decoding,kv_num_headswhen decoding). - decoding (
Bool): Whether the kernel is in the decoding phase. - num_ctas (
UInt32): Number of CTAs to launch (defaults to the H100 SM count). - schedule (
MHASchedule): Strategy for mapping work tiles to thread blocks (defaults toMHASchedule.DEFAULT).
Fields
- gidx_ptr (
Pointer[UInt32, MutAnyOrigin, address_space=AddressSpace.GLOBAL]):
Implemented traits
AnyType,
Copyable,
Deinitable,
DevicePassable,
ImplicitlyCopyable,
MHATileScheduler,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
device_type
comptime device_type = QueuedTileScheduler[tile_shape, num_heads, decoding, num_ctas, schedule]
may_advance
comptime may_advance = True
mha_schedule
comptime mha_schedule = schedule
Methods
__init__
def __init__(gidx_ptr: Pointer[UInt32, MutAnyOrigin]) -> Self
get_current_work_info
def get_current_work_info[ValidLengthType: OptionalPointer, //](self, ts: MHATileSummary[ValidLengthType], state: MHATileState) -> WorkInfo
Returns:
advance
def advance[ValidLengthType: OptionalPointer, //, producer: Bool, sync: MHASchedulerSynchronization = MHASchedulerSynchronization.DEFAULT](self, ts: MHATileSummary[ValidLengthType], mut state: MHATileState, pipeline_idx: UInt32) -> OptionalReg[SeqInfo]
The parameter func must return a Bool indicating whether the WorkInfo arg is valid. This function returns whether the current idx corresponds to a valid WorkInfo. Note that if MHASchedulerSynchronization is NONE, then we assume it is only called by thread_idx.x==0.
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]
Returns:
initial_state
def initial_state[ValidLengthType: OptionalPointer, //](self, ptr: Pointer[UInt32, MutAnyOrigin, address_space=AddressSpace.SHARED], tile_summary: MHATileSummary[ValidLengthType]) -> MHATileState
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).
Returns:
String: The host type's name.