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
gated_group_rmsnorm_gpu
def gated_group_rmsnorm_gpu[dtype: DType, gate_dtype: DType, group_size: Int](output: TileTensor[dtype, Engine=output.Engine, address_space=output.address_space, linear_idx_type=output.linear_idx_type], y: TileTensor[dtype, Engine=y.Engine, address_space=y.address_space, linear_idx_type=y.linear_idx_type], gate: TileTensor[gate_dtype, Engine=gate.Engine, address_space=gate.address_space, linear_idx_type=gate.linear_idx_type], weight: TileTensor[.float32, Engine=weight.Engine, address_space=weight.address_space, linear_idx_type=weight.linear_idx_type], n_rows: Int, num_groups: Int, eps: Float32, ctx: DeviceContext)
Enqueues the fused gated group-RMSNorm; one warp per (row, group).
Parameters:
- dtype (
DType): Element type ofyandoutput. - gate_dtype (
DType): Element type ofgate. - group_size (
Int): Width of each independently normalized group.
Args:
- output (
TileTensor[dtype, Engine=output.Engine, address_space=output.address_space, linear_idx_type=output.linear_idx_type]): The[n_rows, num_groups * group_size]result. - y (
TileTensor[dtype, Engine=y.Engine, address_space=y.address_space, linear_idx_type=y.linear_idx_type]): The[n_rows, num_groups * group_size]SSD scan output. - gate (
TileTensor[gate_dtype, Engine=gate.Engine, address_space=gate.address_space, linear_idx_type=gate.linear_idx_type]): The gate projection; its row stride may exceed the logical width. - weight (
TileTensor[.float32, Engine=weight.Engine, address_space=weight.address_space, linear_idx_type=weight.linear_idx_type]): The fp32 RMSNorm weight. - n_rows (
Int): Number of rows. - num_groups (
Int): Number of groups per row. - eps (
Float32): Epsilon insidersqrt(mean_sq + eps). - ctx (
DeviceContext): Device context to enqueue on.