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

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 the pre output.
  • ​PostLayoutType (TensorLayout): Layout of the post output.
  • ​CombLayoutType (TensorLayout): Layout of the comb output.
  • ​ProjLayoutType (TensorLayout): Layout of the hc_proj input.
  • ​BiasLayoutType (TensorLayout): Layout of the pre_post_comb_b input.
  • ​ScaleLayoutType (TensorLayout): Layout of the pre_post_comb_scale input.
  • ​hc_mult (Int): Number of parallel residual streams.
  • ​hc_sinkhorn_iters (Int): Sinkhorn-Knopp iterations used to project comb.
  • ​warps_per_block (Int): Warps per block.

Args:

Was this page helpful?