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
update_frequency_data_kernel
def update_frequency_data_kernel[freq_data_origin: MutOrigin, FreqDataLayoutType: TensorLayout, freq_offsets_origin: ImmOrigin, FreqOffsetsLayoutType: TensorLayout, new_tokens_origin: ImmOrigin, NewTokensLayoutType: TensorLayout, token_type: DType, block_size: Int](compressed_frequency_data: TileTensor[DType.int32, FreqDataLayoutType, freq_data_origin], frequency_offsets: TileTensor[DType.uint32, FreqOffsetsLayoutType, freq_offsets_origin], new_tokens: TileTensor[token_type, NewTokensLayoutType, new_tokens_origin])
GPU kernel to update token frequency data in CSR format.
Searches for new tokens in existing frequency data and either increments their count or adds them to the first available padding slot.
Parameters:
- freq_data_origin (
MutOrigin): Mutable origin of thecompressed_frequency_datatensor. - FreqDataLayoutType (
TensorLayout): Layout type of thecompressed_frequency_datatensor. - freq_offsets_origin (
ImmOrigin): Immutable origin of thefrequency_offsetstensor. - FreqOffsetsLayoutType (
TensorLayout): Layout type of thefrequency_offsetstensor. - new_tokens_origin (
ImmOrigin): Immutable origin of thenew_tokenstensor. - NewTokensLayoutType (
TensorLayout): Layout type of thenew_tokenstensor. - token_type (
DType): Element type of thenew_tokenstensor. - block_size (
Int): Number of threads per GPU block used to scan a sequence's frequency entries.
Args:
- compressed_frequency_data (
TileTensor[DType.int32, FreqDataLayoutType, freq_data_origin]): 2D CSR frequency data where column 0 is the token id and column 1 is the token count within the sequence, updated in place. - frequency_offsets (
TileTensor[DType.uint32, FreqOffsetsLayoutType, freq_offsets_origin]): 1D tensor of starting indices intocompressed_frequency_datafor each sequence in the batch. - new_tokens (
TileTensor[token_type, NewTokensLayoutType, new_tokens_origin]): 1D tensor of new token ids, one per sequence in the batch.