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
def hyper_connection_gates[hc_mult: Int, hc_sinkhorn_iters: Int, target: StringSpan[ImmStaticOrigin], warps_per_block: Int = Int(4)](pre: TileTensor[.float32, Engine=pre.Engine, address_space=pre.address_space, linear_idx_type=pre.linear_idx_type], post: TileTensor[.float32, Engine=post.Engine, address_space=post.address_space, linear_idx_type=post.linear_idx_type], comb: TileTensor[.float32, Engine=comb.Engine, address_space=comb.address_space, linear_idx_type=comb.linear_idx_type], hc_proj: TileTensor[.float32, Engine=hc_proj.Engine, address_space=hc_proj.address_space, linear_idx_type=hc_proj.linear_idx_type], pre_post_comb_b: TileTensor[.float32, Engine=pre_post_comb_b.Engine, address_space=pre_post_comb_b.address_space, linear_idx_type=pre_post_comb_b.linear_idx_type], pre_post_comb_scale: TileTensor[.float32, Engine=pre_post_comb_scale.Engine, address_space=pre_post_comb_scale.address_space, linear_idx_type=pre_post_comb_scale.linear_idx_type], hc_eps: Float32, context: DeviceContext)
Launches the mHC gate kernel, one warp per row.
Parameters:
- hc_mult (
Int): Number of parallel residual streams. - hc_sinkhorn_iters (
Int): Sinkhorn-Knopp iterations used to projectcomb. - target (
StringSpan[ImmStaticOrigin]): The target device to run the kernel on. - warps_per_block (
Int): Warps per block. Each carriesWARP_SIZE // hc_mult**2rows.
Args:
- pre (
TileTensor[.float32, Engine=pre.Engine, address_space=pre.address_space, linear_idx_type=pre.linear_idx_type]): Stream-collapse weights. Shape:[num_rows, hc_mult]. - post (
TileTensor[.float32, Engine=post.Engine, address_space=post.address_space, linear_idx_type=post.linear_idx_type]): Sublayer-output placement weights. Shape:[num_rows, hc_mult]. - comb (
TileTensor[.float32, Engine=comb.Engine, address_space=comb.address_space, linear_idx_type=comb.linear_idx_type]): Row-major stream mixer. Shape:[num_rows, hc_mult * hc_mult]. - hc_proj (
TileTensor[.float32, Engine=hc_proj.Engine, address_space=hc_proj.address_space, linear_idx_type=hc_proj.linear_idx_type]): Projected streams. Shape:[num_rows, 2 * hc_mult + hc_mult * hc_mult]. - pre_post_comb_b (
TileTensor[.float32, Engine=pre_post_comb_b.Engine, address_space=pre_post_comb_b.address_space, linear_idx_type=pre_post_comb_b.linear_idx_type]): Per-output bias, concatenated inpre,post,comborder. Shape:[2 * hc_mult + hc_mult * hc_mult]. - pre_post_comb_scale (
TileTensor[.float32, Engine=pre_post_comb_scale.Engine, address_space=pre_post_comb_scale.address_space, linear_idx_type=pre_post_comb_scale.linear_idx_type]): Per-output scale, inpre,post,comborder. Shape:[3]. - hc_eps (
Float32): Epsilon guarding the Sinkhorn divisions. - context (
DeviceContext): The device context.
Raises:
If the target is not a GPU or the input widths disagree with hc_mult.