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
hyper_connection_gates_kernel
def hyper_connection_gates_kernel[PreLayoutType: TensorLayout, PostLayoutType: TensorLayout, CombLayoutType: TensorLayout, ProjLayoutType: TensorLayout, BiasLayoutType: TensorLayout, ScaleLayoutType: TensorLayout, hc_mult: Int, hc_sinkhorn_iters: Int, warps_per_block: Int](pre: TileTensor[.float32, PreLayoutType, MutAnyOrigin], post: TileTensor[.float32, PostLayoutType, MutAnyOrigin], comb: TileTensor[.float32, CombLayoutType, MutAnyOrigin], hc_proj: TileTensor[.float32, ProjLayoutType, ImmutAnyOrigin], pre_post_comb_b: TileTensor[.float32, BiasLayoutType, ImmutAnyOrigin], pre_post_comb_scale: TileTensor[.float32, ScaleLayoutType, ImmutAnyOrigin], hc_eps: Float32, num_rows: Int32)
Computes the mHC pre, post and comb gates.
A warp carries WARP_SIZE // hc_mult**2 rows at once, so no lane idles:
two rows per warp at hc_mult=4 on a 32-lane warp, four on a 64-lane one.
Lane l owns row l / hc_mult**2 of the warp's group and, within it,
comb element (t / hc_mult, t % hc_mult) for t = l % hc_mult**2.
That mapping makes both Sinkhorn reductions lane-group reductions of the
same warp: lane_group_sum[num_lanes=hc_mult] sums across a matrix row
(torch's dim=-1) and lane_group_sum[num_lanes=hc_mult, stride=hc_mult] sums down a matrix column (torch's dim=-2). Every xor
mask either reduction uses is smaller than hc_mult**2, so a group never
reaches out of the row that owns it and the rows stay independent.
Parameters:
- PreLayoutType (
TensorLayout): Layout of thepreoutput. - PostLayoutType (
TensorLayout): Layout of thepostoutput. - CombLayoutType (
TensorLayout): Layout of thecomboutput. - ProjLayoutType (
TensorLayout): Layout of thehc_projinput. - BiasLayoutType (
TensorLayout): Layout of thepre_post_comb_binput. - ScaleLayoutType (
TensorLayout): Layout of thepre_post_comb_scaleinput. - hc_mult (
Int): Number of parallel residual streams. - hc_sinkhorn_iters (
Int): Sinkhorn-Knopp iterations used to projectcomb. - warps_per_block (
Int): Warps per block.
Args:
- pre (
TileTensor[.float32, PreLayoutType, MutAnyOrigin]): Stream-collapse weights. Shape:[num_rows, hc_mult]. - post (
TileTensor[.float32, PostLayoutType, MutAnyOrigin]): Sublayer-output placement weights. Shape:[num_rows, hc_mult]. - comb (
TileTensor[.float32, CombLayoutType, MutAnyOrigin]): Row-major stream mixer. Shape:[num_rows, hc_mult * hc_mult]. - hc_proj (
TileTensor[.float32, ProjLayoutType, ImmutAnyOrigin]): Projected streams. Shape:[num_rows, 2 * hc_mult + hc_mult * hc_mult]. - pre_post_comb_b (
TileTensor[.float32, BiasLayoutType, ImmutAnyOrigin]): Per-output bias, concatenated inpre,post,comborder. Shape:[2 * hc_mult + hc_mult * hc_mult]. - pre_post_comb_scale (
TileTensor[.float32, ScaleLayoutType, ImmutAnyOrigin]): Per-output scale, inpre,post,comborder. Shape:[3]. - hc_eps (
Float32): Epsilon guarding the Sinkhorn divisions. - num_rows (
Int32): Number of rows to process.