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
gumbel_sampling_fused_gpu
def gumbel_sampling_fused_gpu[dtype: DType, out_idx_type: DType, //, TemperatureLayoutType: TensorLayout = Layout[TypeList[Int64](), TypeList[ComptimeInt[Int(1)]]()], SeedLayoutType: TensorLayout = Layout[TypeList[Int64](), TypeList[ComptimeInt[Int(1)]]()], from_probs: Bool = False, TemperatureEngine: TensorEngine = DefaultEngine, SeedEngine: TensorEngine = DefaultEngine](ctx: DeviceContext, input: TileTensor[dtype, Engine=input.Engine, address_space=input.address_space, linear_idx_type=input.linear_idx_type], out_idxs: TileTensor[out_idx_type, Engine=out_idxs.Engine, address_space=out_idxs.address_space, linear_idx_type=out_idxs.linear_idx_type], temperature: Optional[TileTensor[.float32, TemperatureLayoutType, ImmutAnyOrigin, Engine=TemperatureEngine]] = None, seed: Optional[TileTensor[.uint64, SeedLayoutType, ImmutAnyOrigin, Engine=SeedEngine]] = None)
Fused Gumbel sampling: applies Gumbel(0,1) noise and selects the argmax in a single GPU kernel launch (no intermediate noised-logits HBM buffer).
Mathematically equivalent to gumbel_sampling_gpu and produces bit-identical
results for the same seed, but saves one full [batch, vocab] HBM
round-trip by fusing noise generation and argmax.
With from_probs the input rows are unnormalized probabilities and the
draw is proportional to them; see _gumbel_argmax_fused_kernel.
Args:
- ctx (
DeviceContext): Device context for GPU operations. - input (
TileTensor[dtype, Engine=input.Engine, address_space=input.address_space, linear_idx_type=input.linear_idx_type]): Input logits tensor [batch, vocab_size]. - out_idxs (
TileTensor[out_idx_type, Engine=out_idxs.Engine, address_space=out_idxs.address_space, linear_idx_type=out_idxs.linear_idx_type]): Output tensor for sampled indices [batch, 1]. - temperature (
Optional[TileTensor[.float32, TemperatureLayoutType, ImmutAnyOrigin, Engine=TemperatureEngine]]): Optional per-token temperature scaling [batch]. - seed (
Optional[TileTensor[.uint64, SeedLayoutType, ImmutAnyOrigin, Engine=SeedEngine]]): Optional per-token random seeds [batch] for reproducibility.