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 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 is 2**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 of logits.
  • ​token_row_offsets (T): Per-batch start offsets into tokens.
  • ​lp_output_offsets (T): Per-batch start offsets into the output row axis.
  • ​lp_output_offsets_host (T): Host-resident copy of lp_output_offsets whose last element gives the total number of output tokens.

Returns:

IndexList[Int(2)]: The output shape [num_output_tokens, 2**levels].

Was this page helpful?