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 function
compute_log_probabilities_ragged_shape
def compute_log_probabilities_ragged_shape[levels: Int](logits: T, tokens: T, sampled_tokens: T, logit_row_offsets: T, token_row_offsets: T, lp_output_offsets: T, lp_output_offsets_host: T) -> IndexList[Int(2)]
Computes the output shapes for the ragged log-probabilities op.
Parameters:
- levels (
Int): Number of heap levels; the output second dimension is2**levels.
Args:
- logits (
T): Input logits ragged by batch, shape[total_rows, vocab_size]. - tokens (
T): Previously generated tokens across all batches, ragged. - sampled_tokens (
T): Most recently sampled token for each batch. - logit_row_offsets (
T): Per-batch start offsets into the first axis oflogits. - token_row_offsets (
T): Per-batch start offsets intotokens. - lp_output_offsets (
T): Per-batch start offsets into the output row axis. - lp_output_offsets_host (
T): Host-resident copy oflp_output_offsetswhose last element gives the total number of output tokens.
Returns:
IndexList[Int(2)]: The output shape [num_output_tokens, 2**levels].