# Copyright (c) 2026 Simulacra Research Inc. # SPDX-License-Identifier: Apache-2.0 from __future__ import annotations import equinox as eqx import jax import jax.numpy as jnp from jaxtyping import Array, Float, Int, PRNGKeyArray from .fused_silu import fused_silu from .readout_leaf_context import ( default_tree_depth, lca_alibi_bias, lca_fixed_slopes, lca_gaussian_decay, lca_gaussian_decay_row, ) from .route_quotient import ( conditional_orbit_ids_from_keys, conditional_orbit_pair_ids_from_keys, ) from .tree import ( CausalRouterEdgeFWLUpdate, EdgeMergeOp, _TREE_NGPT_DEPTH_FEAT_DIM, _tree_ngpt_residual, _tree_clock_root_center_from_depth, _tree_dyadic_segment_clock, _tree_ngpt_level_counts, _tree_sphere, ) QuotientCarrier = tuple[Array, Array] | tuple[Array, Array, Array] def _route_next_pow2(n: int) -> int: return 1 << (int(n) - 1).bit_length() def _dyadic_frontier_add(frontier, value, position): carry = jnp.asarray(value, dtype=frontier.dtype) active = jnp.asarray(True) pos = jnp.asarray(position, dtype=jnp.int32) for level in range(frontier.shape[0]): old = frontier[level] occupied = jnp.bitwise_and(jnp.right_shift(pos, level), 1) == 1 merge = active & occupied place = active & ~occupied frontier = frontier.at[level].set( jnp.where(place, carry, jnp.where(merge, jnp.zeros_like(old), old)) ) carry = jnp.where(merge, old + carry, carry) active = merge return frontier def _dyadic_lca_frontier_sum(frontier, position, w_raw, b): depth = frontier.shape[0] pos = jnp.asarray(position, dtype=jnp.int32) levels = jnp.arange(depth, dtype=jnp.int32) widths = jnp.left_shift(jnp.ones((depth,), dtype=jnp.int32), levels + 1) starts = jnp.bitwise_and(pos, jnp.bitwise_not(widths - 1)) present = jnp.bitwise_and(jnp.right_shift(pos, levels), 1) == 1 decay = lca_gaussian_decay_row(pos, starts, w_raw, b) scale = jnp.where(present[:, None], decay, jnp.zeros_like(decay)) while scale.ndim < frontier.ndim: scale = scale[:, None, :] return jnp.sum(frontier * scale, axis=0) def _replace_square_row_column(matrix, index, row_value, column_value): idx = jnp.arange(matrix.shape[0], dtype=jnp.int32) select = idx == jnp.asarray(index, dtype=jnp.int32) diagonal = row_value[index] column_value = jnp.where(select[:, None], diagonal, column_value) matrix = jnp.where(select[:, None, None], row_value[None, :, :], matrix) return jnp.where( select[None, :, None], column_value[:, None, :], matrix, ) def _square_row_by_reduction(matrix, index): idx = jnp.arange(matrix.shape[0], dtype=jnp.int32) select = idx == jnp.asarray(index, dtype=jnp.int32) return jnp.sum( jnp.where(select[:, None, None], matrix, jnp.zeros_like(matrix)), axis=0, ) def _square_column_local(matrix, index): idx = jnp.arange(matrix.shape[1], dtype=jnp.int32) select = idx == jnp.asarray(index, dtype=jnp.int32) return jnp.sum( jnp.where(select[None, :, None], matrix, jnp.zeros_like(matrix)), axis=1, ) def _route_clock(pos, width: int, dtype, *, base: float | Array = 10000.0, scale=None): if int(width) <= 0: pos_arr = jnp.asarray(pos) return jnp.zeros(pos_arr.shape + (0,), dtype=dtype) pos_f = jnp.asarray(pos, dtype=jnp.float32) if scale is not None: denom = jnp.maximum( jnp.asarray(scale, dtype=jnp.float32) - jnp.asarray(1.0, dtype=jnp.float32), jnp.asarray(1.0, dtype=jnp.float32), ) pos_f = pos_f / denom half = (int(width) + 1) // 2 band = jnp.arange(half, dtype=jnp.float32) base_f = jnp.maximum( jnp.asarray(base, dtype=jnp.float32), jnp.asarray(2.0, dtype=jnp.float32), ) inv_freq = jnp.exp( -jnp.log(base_f) * band / jnp.asarray(max(half, 1), dtype=jnp.float32) ) angle = pos_f[..., None] * inv_freq emb = jnp.concatenate([jnp.sin(angle), jnp.cos(angle)], axis=-1) return emb[..., : int(width)].astype(dtype) def _route_merge_clock( level_idx, pair_idx, pair_base, width: int, max_depth, dtype, *, root_centered: bool = False, ): del pair_base root_center = ( _tree_clock_root_center_from_depth(max_depth, dtype) if root_centered else None ) return _tree_dyadic_segment_clock( level_idx, pair_idx, width, dtype, root_center=root_center, ) class _RoutePointerBase(eqx.Module): w_global: Float[Array, "d_global d_model"] b_global: Float[Array, "d_model"] pref_msg_ln_scale: Float[Array, "two_d_edge"] pref_msg_w1: Float[Array, "two_d_edge d_msg_hidden"] pref_msg_b1: Float[Array, "d_msg_hidden"] pref_msg_w2: Float[Array, "d_msg_hidden d_model"] pref_msg_b2: Float[Array, "d_model"] suff_msg_ln_scale: Float[Array, "two_d_edge"] suff_msg_w1: Float[Array, "two_d_edge d_msg_hidden"] suff_msg_b1: Float[Array, "d_msg_hidden"] suff_msg_w2: Float[Array, "d_msg_hidden d_model"] suff_msg_b2: Float[Array, "d_model"] virt_emb: Float[Array, "one d_model"] order_decay_w: Float[Array, "one d_model"] order_decay_b: Float[Array, "one d_model"] virt_decay_w: Float[Array, "one d_model"] virt_decay_b: Float[Array, "one d_model"] cand_node_ln_scale: Float[Array, "d_in"] cand_global_ln_scale: Float[Array, "d_model"] cand_graw_ln_scale: Float[Array, "d_graw"] cand_g_tap_w: Float[Array, "d_global d_graw"] cand_pref_ln_scale: Float[Array, "d_model"] cand_pref_order_ln_scale: Float[Array, "d_model"] cand_suff_ln_scale: Float[Array, "d_model"] cand_virt_pref_ln_scale: Float[Array, "d_model"] cand_node_w: Float[Array, "d_in d_cand_hidden"] cand_global_w: Float[Array, "d_model d_cand_hidden"] cand_graw_w: Float[Array, "d_graw d_cand_hidden"] cand_pref_w: Float[Array, "d_model d_cand_hidden"] cand_pref_order_w: Float[Array, "d_model d_cand_hidden"] cand_suff_w: Float[Array, "d_model d_cand_hidden"] cand_virt_pref_w: Float[Array, "d_model d_cand_hidden"] cand_virt_ratios_w: Float[Array, "three d_cand_hidden"] cand_b_in: Float[Array, "d_cand_hidden"] cand_block_ln_scale: Float[Array, "b d_cand_hidden"] cand_block_w1: Float[Array, "b d_cand_hidden d_cand_hidden"] cand_block_b1: Float[Array, "b d_cand_hidden"] cand_block_w2: Float[Array, "b d_cand_hidden d_cand_hidden"] cand_block_b2: Float[Array, "b d_cand_hidden"] cand_out_ln_scale: Float[Array, "d_cand_hidden"] cand_w_out: Float[Array, "d_cand_hidden d_model"] cand_b_out: Float[Array, "d_model"] pointer_q_w: Float[Array, "d_model d_score"] pointer_k_w: Float[Array, "d_model d_score"] d_in: int = eqx.field(static=True) d_global: int = eqx.field(static=True) d_edge: int = eqx.field(static=True) d_model: int = eqx.field(static=True) d_attn: int = eqx.field(static=True) pointer_score_dim: int = eqx.field(static=True) n_heads: int = eqx.field(static=True) n_heads_kernel: int = eqx.field(static=True) d_head: int = eqx.field(static=True) max_n: int = eqx.field(static=True) ffn_hidden: int = eqx.field(static=True) msg_hidden: int = eqx.field(static=True) cand_hidden: int = eqx.field(static=True) rope_base: float = eqx.field(static=True) rope_scaling: float = eqx.field(static=True) ln_eps: float = eqx.field(static=True) def __init__( self, *, d_in: int, d_edge: int, d_global: int, d_model: int, n_heads: int, max_n: int, key: PRNGKeyArray, rope_base: float = 10000.0, rope_scaling: float = 1.0, attention_dim: int, pointer_score_dim: int, candidate_hidden: int, summary_hidden: int, ffn_hidden: int, global_tap_dim: int, score_init_scale: float = 1.0, ): ln_eps = 1.0e-5 if d_in != d_model: raise ValueError("router requires d_in == d_model") if d_global < 1: raise ValueError("route pointer d_global must be >= 1") d_attn = int(attention_dim) if d_attn < 1: raise ValueError("route pointer attention_dim must be positive") if d_attn % n_heads != 0: raise ValueError("route pointer attention_dim must be divisible by n_heads") d_head = d_attn // n_heads if d_head % 2 != 0: raise ValueError("route pointer RoPE requires an even per-head dim") score_dim = int(pointer_score_dim) if score_dim < 1: raise ValueError("route pointer pointer_score_dim must be positive") if max_n < 1: raise ValueError("route pointer max_n must be >= 1") if rope_base <= 0.0 or rope_scaling <= 0.0: raise ValueError("route pointer RoPE base/scaling must be positive") n_heads_kernel = 2 * n_heads ffn_hidden = int(ffn_hidden) msg_hidden = int(summary_hidden) if ffn_hidden < 1 or msg_hidden < 1: raise ValueError("route pointer FFN/summary widths must be positive") n_virt_ratios = 3 _graw_dim = int(global_tap_dim) if _graw_dim < 1 or _graw_dim >= int(d_global): raise ValueError( "global_tap_dim must be positive and smaller than d_global" ) cand_in = d_in + 5 * d_model + n_virt_ratios + _graw_dim cand_hidden = int(candidate_hidden) if cand_hidden < 1: raise ValueError("route pointer candidate_hidden must be positive") keys = jax.random.split(key, 22) def w(k, shape, fan_in): return jax.random.normal(k, shape) * (fan_in**-0.5) k_in, k_global = jax.random.split(keys[0], 2) del k_in self.w_global = w(k_global, (int(d_global), d_model), int(d_global)) self.b_global = jnp.zeros((d_model,)) self.cand_graw_ln_scale = jnp.ones((_graw_dim,)) self.cand_g_tap_w = w( jax.random.fold_in(k_global, 0x67AB), (int(d_global), _graw_dim), int(d_global), ) self.pref_msg_ln_scale = jnp.ones((2 * d_edge,)) self.pref_msg_w1 = w(keys[7], (2 * d_edge, msg_hidden), 2 * d_edge) self.pref_msg_b1 = jnp.zeros((msg_hidden,)) self.pref_msg_w2 = w(keys[8], (msg_hidden, d_model), msg_hidden) self.pref_msg_b2 = jnp.zeros((d_model,)) self.suff_msg_ln_scale = jnp.ones((2 * d_edge,)) self.suff_msg_w1 = w(keys[9], (2 * d_edge, msg_hidden), 2 * d_edge) self.suff_msg_b1 = jnp.zeros((msg_hidden,)) self.suff_msg_w2 = w(keys[10], (msg_hidden, d_model), msg_hidden) self.suff_msg_b2 = jnp.zeros((d_model,)) vkey = jax.random.fold_in(key, 0x5710C) self.virt_emb = jax.random.normal(vkey, (1, d_model)) * (d_model**-0.5) from .readout_leaf_context import lca_order_init_w_b self.order_decay_w, self.order_decay_b = lca_order_init_w_b(d_model) self.virt_decay_w, self.virt_decay_b = lca_order_init_w_b(d_model) self.cand_node_ln_scale = jnp.ones((d_in,)) self.cand_global_ln_scale = jnp.ones((d_model,)) self.cand_pref_ln_scale = jnp.ones((d_model,)) self.cand_pref_order_ln_scale = jnp.ones((d_model,)) self.cand_suff_ln_scale = jnp.ones((d_model,)) self.cand_virt_pref_ln_scale = jnp.ones((d_model,)) compose_key = keys[15] self.cand_node_w = w( jax.random.fold_in(compose_key, 0), (d_in, cand_hidden), cand_in, ) self.cand_global_w = w( jax.random.fold_in(compose_key, 1), (d_model, cand_hidden), cand_in, ) self.cand_graw_w = w( jax.random.fold_in(compose_key, 2), (_graw_dim, cand_hidden), cand_in, ) self.cand_pref_w = w( jax.random.fold_in(compose_key, 3), (d_model, cand_hidden), cand_in, ) self.cand_pref_order_w = w( jax.random.fold_in(compose_key, 4), (d_model, cand_hidden), cand_in, ) self.cand_suff_w = w( jax.random.fold_in(compose_key, 5), (d_model, cand_hidden), cand_in, ) self.cand_virt_pref_w = w( jax.random.fold_in(compose_key, 6), (d_model, cand_hidden), cand_in, ) self.cand_virt_ratios_w = w( jax.random.fold_in(compose_key, 10), (n_virt_ratios, cand_hidden), cand_in, ) self.cand_b_in = jnp.zeros((cand_hidden,)) self.cand_block_ln_scale = jnp.ones((1, cand_hidden)) self.cand_block_w1 = w( keys[16], (1, cand_hidden, cand_hidden), cand_hidden, ) self.cand_block_b1 = jnp.zeros((1, cand_hidden)) self.cand_block_w2 = w( keys[17], (1, cand_hidden, cand_hidden), cand_hidden, ) self.cand_block_b2 = jnp.zeros((1, cand_hidden)) cand_out_key = keys[18] self.cand_out_ln_scale = jnp.ones((cand_hidden,)) self.cand_w_out = w(cand_out_key, (cand_hidden, d_model), cand_hidden) self.cand_b_out = jnp.zeros((d_model,)) pointer_q_key = keys[19] pointer_k_key = keys[20] self.pointer_q_w = w(pointer_q_key, (d_model, score_dim), d_model) * float( score_init_scale ) self.pointer_k_w = w(pointer_k_key, (d_model, score_dim), d_model) del keys self.d_in = d_in self.d_global = int(d_global) self.d_edge = d_edge self.d_model = d_model self.d_attn = d_attn self.pointer_score_dim = score_dim self.n_heads = n_heads self.n_heads_kernel = n_heads_kernel self.d_head = d_head self.max_n = max_n self.ffn_hidden = ffn_hidden self.msg_hidden = msg_hidden self.cand_hidden = cand_hidden self.rope_base = float(rope_base) self.rope_scaling = float(rope_scaling) self.ln_eps = float(ln_eps) def _ln( self, scale, x, *, tag_id: str, kfac_structural_mask=None, kfac_scan_shared: bool = False, kfac_repeat_ndim: int = 0, kfac_context_primal_reused_over_walkers: bool = False, ): from hamiltonzero.model.tree import _tagged_rms_eqx_style return _tagged_rms_eqx_style( scale, x, eps=self.ln_eps, tag_id=tag_id, pathway="even", var_floor=1e-2, kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) def _cross_ln( self, scale, shift, x, *, tag_id: str, kfac_structural_mask=None, kfac_scan_shared: bool = False, kfac_repeat_ndim: int = 0, kfac_context_primal_reused_over_walkers: bool = False, ): from hamiltonzero.model.tree import _tagged_ln_eqx_style return _tagged_ln_eqx_style( scale, shift, x, eps=self.ln_eps, tag_id=tag_id, pathway="even", var_floor=1e-2, kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) def _dense( self, weight, bias, x, *, tag_id: str, kfac_structural_mask=None, kfac_scan_shared: bool = False, kfac_repeat_ndim: int = 0, kfac_context_primal_reused_over_walkers: bool = False, ): from hamiltonzero.model.tree import _tagged_dense return _tagged_dense( weight, bias, x, tag_id=tag_id, pathway="even", kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) def _dense_no_bias( self, weight, x, *, tag_id: str, kfac_structural_mask=None, kfac_scan_shared: bool = False, kfac_repeat_ndim: int = 0, kfac_context_primal_reused_over_walkers: bool = False, ): from hamiltonzero.model.tree import _tagged_dense_no_bias return _tagged_dense_no_bias( weight, x, tag_id=tag_id, pathway="even", kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) def _project_nodes(self, h: Float[Array, "n d_in"], structural_mask=None): del structural_mask return h def _prepare_nodes( self, h: Float[Array, "n d_in"], mask: Int[Array, "n"] | Array, ): projected, node_mean = self._center_nodes( self._project_nodes(h, mask.astype(bool)), mask, ) return (h, projected), node_mean def _project_global( self, global_feat: Float[Array, "d_global"], dtype, structural_mask=None, ): raw = global_feat.astype(dtype) g_dm = self._dense( self.w_global, self.b_global, raw, tag_id="route.global_input", kfac_structural_mask=structural_mask, kfac_repeat_ndim=0, ) return (raw, g_dm) def _center_nodes( self, node_state: Float[Array, "n d_model"], mask: Int[Array, "n"] | Array, ): dtype = node_state.dtype active = mask.astype(dtype).reshape(node_state.shape[0], 1) denom = jnp.maximum(jnp.sum(active), jnp.asarray(1.0, dtype=dtype)) global_state = jnp.sum(node_state * active, axis=0) / denom return node_state, global_state def _message_mlp(self, edge_pair, *, prefix: bool, structural_mask=None): if prefix: ln_s = self.pref_msg_ln_scale w1, b1, w2, b2 = ( self.pref_msg_w1, self.pref_msg_b1, self.pref_msg_w2, self.pref_msg_b2, ) name = "pref" else: ln_s = self.suff_msg_ln_scale w1, b1, w2, b2 = ( self.suff_msg_w1, self.suff_msg_b1, self.suff_msg_w2, self.suff_msg_b2, ) name = "suff" structural_mask = ( jnp.ones(edge_pair.shape[:-1], dtype=bool) if structural_mask is None else jnp.broadcast_to( jnp.asarray(structural_mask, dtype=bool), edge_pair.shape[:-1] ) ) kfac_kwargs = dict( kfac_structural_mask=structural_mask, kfac_repeat_ndim=structural_mask.ndim, ) x = self._ln( ln_s, edge_pair, tag_id=f"route.candidate.{name}_msg_ln", **kfac_kwargs, ) x = self._dense( w1, b1, x, tag_id=f"route.candidate.{name}_msg1", **kfac_kwargs, ) x = fused_silu(x) return self._dense( w2, b2, x, tag_id=f"route.candidate.{name}_msg2", **kfac_kwargs, ) def _edge_pair_for_source( self, edge: Float[Array, "n n d_edge"], source: Int[Array, ""], ) -> Float[Array, "n two_d_edge"]: return jnp.concatenate( [edge[:, source, :], edge[source, :, :]], axis=-1, ) def _ordered_edge_messages( self, edge: Float[Array, "n n d_edge"], perm: Int[Array, "n"], mask: Int[Array, "n"] | Array, ): n = edge.shape[0] idx = jnp.arange(n, dtype=jnp.int32) edge_i_p = edge[idx[None, :], perm[:, None], :] edge_p_i = edge[perm[:, None], idx[None, :], :] edge_pair = jnp.concatenate([edge_i_p, edge_p_i], axis=-1) mask_bool = mask.astype(bool) pair_structural_mask = mask_bool[perm][:, None] & mask_bool[None, :] return ( self._message_mlp( edge_pair, prefix=True, structural_mask=pair_structural_mask, ), self._message_mlp( edge_pair, prefix=False, structural_mask=pair_structural_mask, ), ) def _clock_root_center_from_mask(self, mask): n_active = jnp.maximum(jnp.sum(jnp.asarray(mask, dtype=jnp.int32)), 1) depth = jnp.ceil(jnp.log2(n_active.astype(jnp.float32))).astype(jnp.int32) return _tree_clock_root_center_from_depth(depth, jnp.float32) def _route_position_embedding(self, pos, dtype, *, mask=None): pos_f = jnp.asarray(pos, dtype=jnp.float32) if mask is not None: pos_f = pos_f - self._clock_root_center_from_mask(mask) pos_f = pos_f / jnp.asarray( self.rope_scaling, dtype=jnp.float32, ) half = (self.d_model + 1) // 2 band = jnp.arange(half, dtype=jnp.float32) inv_freq = jnp.exp( -jnp.log(jnp.asarray(self.rope_base, dtype=jnp.float32)) * band / jnp.asarray(max(half, 1), dtype=jnp.float32) ) angle = pos_f[..., None] * inv_freq emb = jnp.concatenate([jnp.sin(angle), jnp.cos(angle)], axis=-1) return emb[..., : self.d_model].astype(dtype) def _first_active_index(self, mask): return jnp.argmax(mask.astype(jnp.int32)).astype(jnp.int32) def _compose_candidates( self, node_state, global_state: Float[Array, "d_global"], prefix_summary: Float[Array, "... n d_model"], prefix_order_summary: Float[Array, "... n d_model"], suffix_summary: Float[Array, "... n d_model"], route_pos, virt_pref_order_summary: Float[Array, "... n d_model"], virt_ratios: Float[Array, "... n 3"], clock_mask=None, candidate_mask=None, ) -> Float[Array, "... n d_model"]: node_input, node_projected = node_state g_raw, g_dm = global_state candidate_structural_mask = ( jnp.ones(prefix_summary.shape[:-1], dtype=bool) if candidate_mask is None else jnp.broadcast_to( jnp.asarray(candidate_mask, dtype=bool), prefix_summary.shape[:-1], ) ) kfac_kwargs = dict( kfac_structural_mask=candidate_structural_mask, kfac_repeat_ndim=candidate_structural_mask.ndim, ) from hamiltonzero.model.tree import _tagged_dense_no_bias g_raw = _tagged_dense_no_bias( self.cand_g_tap_w, g_raw, tag_id="route.candidate.gtap", pathway="even", kfac_structural_mask=jnp.any(candidate_structural_mask), kfac_repeat_ndim=0, ) if prefix_summary.ndim == node_projected.ndim: nodes = node_projected node_inputs = node_input global_nodes = jnp.broadcast_to(g_dm[None, :], nodes.shape) graw_nodes = jnp.broadcast_to( g_raw[None, :], nodes.shape[:-1] + (g_raw.shape[-1],) ) else: nodes = jnp.broadcast_to( node_projected, prefix_summary.shape[:-1] + (self.d_model,), ) node_inputs = jnp.broadcast_to( node_input, prefix_summary.shape[:-1] + (self.d_in,), ) global_nodes = jnp.broadcast_to( g_dm, prefix_summary.shape[:-1] + (self.d_model,), ) graw_nodes = jnp.broadcast_to( g_raw, prefix_summary.shape[:-1] + (g_raw.shape[-1],), ) pos_nodes = self._route_position_embedding( route_pos, prefix_summary.dtype, mask=clock_mask, ) while pos_nodes.ndim < global_nodes.ndim: pos_nodes = pos_nodes[..., None, :] global_nodes = global_nodes + jnp.broadcast_to(pos_nodes, global_nodes.shape) node_in = self._ln( self.cand_node_ln_scale, node_inputs, tag_id="route.candidate.node_ln", **kfac_kwargs, ) global_in = self._ln( self.cand_global_ln_scale, global_nodes, tag_id="route.candidate.global_ln", **kfac_kwargs, ) graw_in = self._ln( self.cand_graw_ln_scale, graw_nodes, tag_id="route.candidate.graw_ln", **kfac_kwargs, ) pref_in = self._ln( self.cand_pref_ln_scale, prefix_summary, tag_id="route.candidate.pref_ln", **kfac_kwargs, ) pref_order_in = self._ln( self.cand_pref_order_ln_scale, prefix_order_summary, tag_id="route.candidate.pref_order_ln", **kfac_kwargs, ) suff_in = self._ln( self.cand_suff_ln_scale, suffix_summary, tag_id="route.candidate.suff_ln", **kfac_kwargs, ) _vp = jnp.broadcast_to(virt_pref_order_summary, prefix_summary.shape) virt_pref_in = self._ln( self.cand_virt_pref_ln_scale, _vp, tag_id="route.candidate.virt_pref_ln", **kfac_kwargs, ) virt_ratios_in = jnp.broadcast_to( virt_ratios, prefix_summary.shape[:-1] + (3,), ).astype(prefix_summary.dtype) x = self._dense( self.cand_node_w, self.cand_b_in, node_in, tag_id="route.candidate.compose.node", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_global_w, global_in, tag_id="route.candidate.compose.global", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_graw_w, graw_in, tag_id="route.candidate.compose.graw", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_pref_w, pref_in, tag_id="route.candidate.compose.pref", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_pref_order_w, pref_order_in, tag_id="route.candidate.compose.pref_order", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_suff_w, suff_in, tag_id="route.candidate.compose.suff", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_virt_pref_w, virt_pref_in, tag_id="route.candidate.compose.virt_pref", **kfac_kwargs, ) x = x + self._dense_no_bias( self.cand_virt_ratios_w, virt_ratios_in, tag_id="route.candidate.compose.virt_ratios", **kfac_kwargs, ) def block_body(x, params): ln_s, w1, b1, w2, b2 = params y = self._ln( ln_s, x, tag_id="route.candidate.block.ln", **kfac_kwargs, ) y = self._dense( w1, b1, y, tag_id="route.candidate.block.ffn1", **kfac_kwargs, ) y = fused_silu(y) y = self._dense( w2, b2, y, tag_id="route.candidate.block.ffn2", **kfac_kwargs, ) return x + y, None x, _ = jax.lax.scan( block_body, x, ( self.cand_block_ln_scale, self.cand_block_w1, self.cand_block_b1, self.cand_block_w2, self.cand_block_b2, ), ) x = self._ln( self.cand_out_ln_scale, x, tag_id="route.candidate.compose_out_ln", **kfac_kwargs, ) delta = self._dense( self.cand_w_out, self.cand_b_out, x, tag_id="route.candidate.compose_out", **kfac_kwargs, ) return nodes + delta def _teacher_candidate_states( self, node_state, global_state: Float[Array, "d_global"], edge: Float[Array, "n n d_edge"], perm: Int[Array, "n"], mask: Int[Array, "n"] | Array, real_mask: Int[Array, "n"] | Array, ) -> Float[Array, "n n d_model"]: _node_input, node_projected = node_state n = node_projected.shape[0] dtype = node_projected.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) rm_bool = real_mask.astype(bool) virt_slot = mask_bool & (~rm_bool) virt_at_pos = virt_slot[perm].astype(dtype) pref_msg, suff_msg = self._ordered_edge_messages(edge, perm, mask) row_active = mask_bool.astype(dtype).reshape(n, 1, 1) pref_msg = pref_msg * row_active suff_msg = suff_msg * row_active pref_before = jnp.cumsum(pref_msg, axis=0) - pref_msg from .readout_leaf_context import lca_gaussian_decay, register_vector_as_dense _odw = register_vector_as_dense( self.order_decay_w, tag_id="route.order_decay_w" )[0] _odb = register_vector_as_dense( self.order_decay_b, tag_id="route.order_decay_b" )[0] _tri = (idx[:, None] > idx[None, :]).astype(dtype) _odecay = lca_gaussian_decay(idx, idx, _odw, _odb) pref_order_before = jnp.einsum("ts,tsd,sid->tid", _tri, _odecay, pref_msg) suff_before = jnp.cumsum(suff_msg, axis=0) - suff_msg suff_including_self = jnp.sum(suff_msg, axis=0)[None, :, :] - suff_before pos_of_node = jnp.zeros((n,), dtype=jnp.int32).at[perm].set(idx) self_msg = suff_msg[pos_of_node, idx, :] remaining = mask_bool[None, :] & (pos_of_node[None, :] >= idx[:, None]) candidate_structural_mask = mask_bool[:, None] & remaining suff_other = suff_including_self - jnp.where( remaining[:, :, None], self_msg[None, :, :], 0.0, ) virt_emb = register_vector_as_dense( self.virt_emb, tag_id="route.virt_emb", )[0] virt_msg = virt_at_pos[:, None] * virt_emb[None, :] _vdw = register_vector_as_dense(self.virt_decay_w, tag_id="route.virt_decay_w")[ 0 ] _vdb = register_vector_as_dense(self.virt_decay_b, tag_id="route.virt_decay_b")[ 0 ] _vdecay = lca_gaussian_decay(idx, idx, _vdw, _vdb) virt_pref_order_before = jnp.einsum( "ts,tsd,sd->td", _tri, _vdecay, virt_msg, ) virt_cnt_prefix = jnp.cumsum(virt_at_pos) - virt_at_pos total_empty = jnp.sum(virt_at_pos) total_leafs = jnp.sum(mask_bool.astype(dtype)) virt_cnt_suffix = total_empty - virt_cnt_prefix virt_norm = jnp.sqrt(jnp.maximum(virt_cnt_prefix, 1.0))[:, None] virt_ratios = jnp.stack( [ virt_cnt_suffix / jnp.maximum(total_empty, 1.0), virt_cnt_suffix / jnp.maximum(total_leafs, 1.0), jnp.log((virt_cnt_prefix + 1.0) / (virt_cnt_suffix + 1.0)), ], axis=-1, ) virt_pref_order_summary = (virt_pref_order_before / virt_norm)[:, None, :] virt_ratios_summary = virt_ratios[:, None, :] pref_den = jnp.sqrt(jnp.maximum(idx, 1).astype(dtype)).reshape(n, 1, 1) n_active = jnp.sum(mask_bool.astype(jnp.int32)) suff_den = jnp.sqrt(jnp.maximum(n_active - idx - 1, 1).astype(dtype)).reshape( n, 1, 1 ) return self._compose_candidates( node_state, global_state, pref_before / pref_den, pref_order_before / pref_den, suff_other / suff_den, idx, virt_pref_order_summary=virt_pref_order_summary, virt_ratios=virt_ratios_summary, clock_mask=mask, candidate_mask=candidate_structural_mask, ) def _initial_summaries_and_edge_messages( self, edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, dtype, ): n = edge.shape[0] idx = jnp.arange(n, dtype=jnp.int32) pref_msg, suff_msg = self._ordered_edge_messages(edge, idx, mask) source_active = mask.astype(dtype).reshape(n, 1, 1) not_self = (idx[:, None] != idx[None, :]).astype(dtype).reshape(n, n, 1) suffix_raw = jnp.sum(suff_msg * source_active * not_self, axis=0) zeros = jnp.zeros_like(suffix_raw) virt_prefix_order0 = jnp.zeros((self.d_model,), dtype=suffix_raw.dtype) virt_count0 = jnp.zeros((), dtype=suffix_raw.dtype) summaries = ( zeros, zeros, suffix_raw, virt_prefix_order0, virt_count0, ) return summaries, (pref_msg, suff_msg) def _initial_summaries( self, edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, dtype, ): summaries, _edge_messages = self._initial_summaries_and_edge_messages( edge, mask, dtype ) return summaries def _initial_summaries_streamed( self, edge: Float[Array, "n n d_edge"], edge_transpose: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, dtype, *, pair_tile_size: int | None = None, sequence_axis_name: str | None = None, sequence_mesh=None, ): n = edge.shape[0] idx = jnp.arange(n, dtype=jnp.int32) def _seq_constraint(value, *axes): if sequence_axis_name is None: return value from jax.sharding import NamedSharding, PartitionSpec as P spec = P(*axes) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) tile = n if pair_tile_size is None else min(int(pair_tile_size), n) if tile < 1: raise ValueError("pair_tile_size must be positive") n_tiles = (n + tile - 1) // tile padded_n = n_tiles * tile source_pad = padded_n - n edge_padded = jnp.pad(edge, ((0, 0), (0, source_pad), (0, 0))) edge_transpose_padded = jnp.pad( edge_transpose, ((0, 0), (0, source_pad), (0, 0)), ) source_mask = jnp.pad(mask.astype(bool), ((0, source_pad),)) candidate_mask = mask.astype(bool)[:, None] candidate_ids = idx[:, None] suffix0 = _seq_constraint( jnp.zeros((n, self.d_model), dtype=dtype), sequence_axis_name, None, ) def add_source_tile(tile_index, suffix_sum): start = tile_index * tile edge_tile = jax.lax.dynamic_slice_in_dim( edge_padded, start, tile, axis=1, ) edge_transpose_tile = jax.lax.dynamic_slice_in_dim( edge_transpose_padded, start, tile, axis=1, ) edge_pair = jnp.concatenate( [edge_tile, edge_transpose_tile], axis=-1, ) source_mask_tile = jax.lax.dynamic_slice_in_dim( source_mask, start, tile, axis=0, ) source_ids = start + jnp.arange(tile, dtype=jnp.int32) pair_mask = candidate_mask & source_mask_tile[None, :] suff_by_candidate = self._message_mlp( edge_pair, prefix=False, structural_mask=pair_mask, ) source_weight = source_mask_tile.astype(dtype)[None, :, None] not_self = (candidate_ids != source_ids[None, :]).astype(dtype) suffix_sum = suffix_sum + jnp.sum( suff_by_candidate * source_weight * not_self[..., None], axis=1, ) return _seq_constraint( suffix_sum, sequence_axis_name, None, ) suffix_raw = jax.lax.fori_loop(0, n_tiles, add_source_tile, suffix0) zeros = jnp.zeros_like(suffix_raw) return ( zeros, zeros, suffix_raw, jnp.zeros((self.d_model,), dtype=suffix_raw.dtype), jnp.zeros((), dtype=suffix_raw.dtype), ) def _candidate_states_from_summaries( self, node_state, global_state: Float[Array, "d_global"], prefix_raw: Float[Array, "n d_model"], prefix_order_raw: Float[Array, "n d_model"], suffix_raw: Float[Array, "n d_model"], route_pos: Int[Array, ""] | Float[Array, ""], edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, prefix_ids: Int[Array, "n"], virt_prefix_order_raw: Float[Array, "d_model"], virt_count: Float[Array, ""], real_mask: Int[Array, "n"] | Array, ) -> Float[Array, "n d_model"]: _node_input, node_projected = node_state dtype = node_projected.dtype route_pos_i = jnp.asarray(route_pos, dtype=jnp.int32) mask_bool = mask.astype(bool) pref_den = jnp.sqrt(jnp.maximum(route_pos_i, 1).astype(dtype)) n_active = jnp.sum(mask_bool.astype(jnp.int32)) suff_den = jnp.sqrt(jnp.maximum(n_active - route_pos_i - 1, 1).astype(dtype)) _cand_idx = jnp.arange(node_projected.shape[0], dtype=jnp.int32) _placed_pos = jnp.arange(node_projected.shape[0], dtype=jnp.int32) _already_picked = jnp.any( (_placed_pos < route_pos_i)[:, None] & (prefix_ids[:, None] == _cand_idx[None, :]), axis=0, ) candidate_structural_mask = mask_bool & ~_already_picked _vnorm = jnp.sqrt(jnp.maximum(virt_count, 1.0)) _vp = virt_prefix_order_raw / _vnorm virt_slot = mask_bool & (~real_mask.astype(bool)) total_empty = jnp.sum(virt_slot.astype(dtype)) total_leafs = jnp.sum(mask.astype(dtype)) virt_cnt_suffix = total_empty - virt_count _vr = jnp.stack( [ virt_cnt_suffix / jnp.maximum(total_empty, 1.0), virt_cnt_suffix / jnp.maximum(total_leafs, 1.0), jnp.log((virt_count + 1.0) / (virt_cnt_suffix + 1.0)), ], ) return self._compose_candidates( node_state, global_state, prefix_raw / pref_den, prefix_order_raw / pref_den, suffix_raw / suff_den, route_pos_i, virt_pref_order_summary=_vp, virt_ratios=_vr, clock_mask=mask, candidate_mask=candidate_structural_mask, ) def _pointer_raw(self, hidden, candidate_state, structural_mask=None): candidate_structural_mask = ( jnp.ones(candidate_state.shape[:-1], dtype=bool) if structural_mask is None else jnp.broadcast_to( jnp.asarray(structural_mask, dtype=bool), candidate_state.shape[:-1], ) ) if hidden.ndim == 1: query_structural_mask = jnp.any(candidate_structural_mask) q = self._dense_no_bias( self.pointer_q_w, hidden, tag_id="route.pointer.q", kfac_structural_mask=query_structural_mask, kfac_repeat_ndim=0, ) k = self._dense_no_bias( self.pointer_k_w, candidate_state, tag_id="route.pointer.k", kfac_structural_mask=candidate_structural_mask, kfac_repeat_ndim=1, ) raw = jnp.einsum("d,nd->n", q, k) else: query_structural_mask = jnp.any(candidate_structural_mask, axis=-1) q = self._dense_no_bias( self.pointer_q_w, hidden, tag_id="route.pointer.q", kfac_structural_mask=query_structural_mask, kfac_repeat_ndim=1, ) k = self._dense_no_bias( self.pointer_k_w, candidate_state, tag_id="route.pointer.k", kfac_structural_mask=candidate_structural_mask, kfac_repeat_ndim=2, ) raw = jnp.einsum("td,tnd->tn", q, k) scale = jax.lax.rsqrt(jnp.asarray(self.pointer_score_dim, dtype=raw.dtype)) return raw * scale def _pointer_logits(self, hidden, candidate_state, picked, mask, tau): active = mask.astype(bool) & (~picked) raw = self._pointer_raw(hidden, candidate_state, structural_mask=active) neg = jnp.asarray(-1.0e30, dtype=raw.dtype) return jnp.where(active, raw / jnp.asarray(tau, dtype=raw.dtype), neg) def _learned_first_choice_mask(self, mask, real_mask): mask_bool = mask.astype(bool) if real_mask is None: return mask_bool real_bool = real_mask.astype(bool) & mask_bool return jnp.where(jnp.any(real_bool), real_bool, mask_bool) def _step_choice_mask(self, first_step, mask, real_mask): mask_bool = mask.astype(bool) first_mask = self._learned_first_choice_mask(mask, real_mask) return jnp.where(first_step, first_mask, mask_bool) def _step_pointer_hidden(self, first_step, global_state, hidden): first_hidden = jnp.broadcast_to(global_state[1], hidden.shape) return jnp.where(first_step, first_hidden, hidden) def _teacher_hidden_with_first(self, hidden, global_state, first_active_idx): return hidden.at[first_active_idx].set(global_state[1]) def _score_step_for_logp(self, first_step, predict_step): return predict_step | first_step def _logprob_contribute_mask(self, mask_bool, idx, first_active): del idx, first_active return mask_bool def _collapse_quotient_logits(self, logits, ids, valid): n = logits.shape[-1] dtype = logits.dtype idx = jnp.arange(n, dtype=jnp.int32) ids = jnp.asarray(ids, dtype=jnp.int32) valid = valid.astype(bool) & (ids >= 0) same = (ids[:, None] == ids[None, :]) & valid[:, None] & valid[None, :] rep_idx = jnp.min( jnp.where(same, idx[None, :], jnp.asarray(n, dtype=jnp.int32)), axis=1, ) reps = valid & (idx == rep_idx) neg = jnp.asarray(-1.0e30, dtype=dtype) member_logits = jnp.where(same, logits[None, :], neg) max_l = jnp.max(member_logits, axis=1) max_l = jnp.where(jnp.isfinite(max_l), max_l, jnp.asarray(0.0, dtype=dtype)) class_lse = max_l + jnp.log( jnp.sum(jnp.exp(member_logits - max_l[:, None]), axis=1) ) class_size = jnp.maximum( jnp.sum(same.astype(dtype), axis=1), jnp.asarray(1.0, dtype=dtype), ) quotient_logits = class_lse - jnp.log(class_size) return jnp.where(reps, quotient_logits, neg) def _apply_quotient_logits( self, logits, first_orbit_ids, valid_mask, context_mask, prefix_ids, prefix_len, ): if len(first_orbit_ids) == 2: node_key, edge_key = first_orbit_ids ids = conditional_orbit_ids_from_keys( node_key, edge_key, valid_mask, context_mask, prefix_ids, prefix_len, ) elif len(first_orbit_ids) == 3: node_key, edge_key, needs_fwl2 = first_orbit_ids inputs = ( node_key, edge_key, valid_mask, context_mask, prefix_ids, prefix_len, ) ids = jax.lax.cond( jnp.asarray(needs_fwl2, dtype=jnp.bool_), lambda values: conditional_orbit_pair_ids_from_keys(*values), lambda values: conditional_orbit_ids_from_keys(*values), inputs, ) else: raise ValueError( "quotient carrier must contain node key, edge key, and " "optionally needs_fwl2" ) return self._collapse_quotient_logits(logits, ids, valid_mask) def _stopgrad_logit_scale(self, logits, valid_mask): del valid_mask return logits def _append_token( self, token: Float[Array, "d_model"], chosen: Int[Array, ""], t: Int[Array, ""], prefix_ids: Int[Array, "n"], k_cache: Float[Array, "l n h d_head"], v_cache: Float[Array, "l n h d_head"], edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, ): del chosen, t, prefix_ids, edge, mask return token, k_cache, v_cache def logprob_perm( self, h: Float[Array, "n d_in"], edge: Float[Array, "n n d_edge"], perm: Int[Array, "n"], mask: Int[Array, "n"] | Array, *, global_feat: Float[Array, "d_global"] | None = None, tau: float | Float[Array, ""] = 1.0, real_mask: Int[Array, "n"] | Array | None = None, first_orbit_ids: QuotientCarrier, ) -> Float[Array, ""]: scores = self._teacher_logits( h, edge, perm, mask, global_feat=global_feat, tau=tau, real_mask=real_mask, first_orbit_ids=first_orbit_ids, ) n = h.shape[0] if n <= 1: return jnp.asarray(0.0, dtype=h.dtype) mask_bool = mask.astype(bool) first_active = self._first_active_index(mask) idx = jnp.arange(n, dtype=jnp.int32) contribute = self._logprob_contribute_mask(mask_bool, idx, first_active) neg = jnp.asarray(-1.0e30, dtype=scores.dtype) scores = self._stopgrad_logit_scale( scores, scores > (neg * jnp.asarray(0.5, dtype=scores.dtype)), ) log_probs = jax.nn.log_softmax(scores.astype(jnp.float32), axis=-1) chosen = jnp.take_along_axis(log_probs, perm[:, None], axis=-1)[:, 0] return jnp.sum(jnp.where(contribute, chosen, 0.0)) def logprob_identity( self, h: Float[Array, "n d_in"], edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, *, global_feat: Float[Array, "d_global"] | None = None, tau: float | Float[Array, ""] = 1.0, real_mask: Int[Array, "n"] | Array | None = None, first_orbit_ids: QuotientCarrier, ) -> Float[Array, ""]: n = h.shape[0] return self.logprob_perm( h, edge, jnp.arange(n, dtype=jnp.int32), mask, global_feat=global_feat, tau=tau, real_mask=real_mask, first_orbit_ids=first_orbit_ids, ) class _TreePrefixMerge(eqx.Module): ln_scale: Float[Array, "d_in"] g_proj_w: Float[Array, "d_gstream d_gsec"] w1: Float[Array, "d_in d_hidden"] b1: Float[Array, "d_hidden"] w2: Float[Array, "d_hidden d_model"] b2: Float[Array, "d_model"] alpha_route: Float[Array, "d_model"] d_model: int = eqx.field(static=True) d_hidden: int = eqx.field(static=True) d_in: int = eqx.field(static=True) max_depth: int = eqx.field(static=True) ngpt_alpha_max: float = eqx.field(static=True) ln_eps: float = eqx.field(static=True) def __init__( self, d_model: int, *, hidden: int, max_depth: int, key: PRNGKeyArray, gladder_d_g: int, alpha_init: float, alpha_max: float, ln_eps: float = 1e-5, ): d_hidden = int(hidden) d_in = 5 * int(d_model) + 64 + _TREE_NGPT_DEPTH_FEAT_DIM k1, k2 = jax.random.split(key, 2) self.ln_scale = jnp.ones((d_in,)) self.w1 = jax.random.normal(k1, (d_in, d_hidden)) * (d_in**-0.5) self.b1 = jnp.zeros((d_hidden,)) self.w2 = jax.random.normal(k2, (d_hidden, d_model)) * (d_hidden**-0.5) self.b2 = jnp.zeros((d_model,)) kg = jax.random.fold_in(k2, 0x61B5) self.g_proj_w = jax.random.normal(kg, (int(gladder_d_g), 64)) * ( int(gladder_d_g) ** -0.5 ) self.alpha_route = float(alpha_init) * jnp.ones((int(d_model),)) self.d_model = int(d_model) self.d_hidden = int(d_hidden) self.d_in = int(d_in) self.max_depth = int(max_depth) self.ngpt_alpha_max = float(alpha_max) self.ln_eps = float(ln_eps) def project_global(self, g, structural_mask): from hamiltonzero.model.tree import _tagged_dense_no_bias return _tagged_dense_no_bias( self.g_proj_w, g, tag_id="gladder.route.merge_gproj", pathway="even", kfac_structural_mask=jnp.asarray(structural_mask, dtype=bool), kfac_scan_shared=False, kfac_repeat_ndim=0, kfac_context_primal_reused_over_walkers=True, ) def __call__( self, left, right, left_mask, right_mask, sibling_edge_lr, sibling_edge_rl, level_idx, pair_idx=None, pair_base=None, clock_depth=None, depth_feats=None, g=None, g_structural_mask=None, g_projected=None, ): from hamiltonzero.model.tree import _tagged_dense, _tagged_rms_eqx_style from hamiltonzero.model.tree import _rownorm_cols _we = _rownorm_cols dtype = left.dtype left_mask = left_mask.astype(dtype) right_mask = right_mask.astype(dtype) out_mask = left_mask + right_mask - left_mask * right_mask both = left_mask * right_mask merge_structural_mask = both.astype(bool) depth_i = ( jnp.asarray(max(self.max_depth, 1), dtype=jnp.int32) if clock_depth is None else jnp.maximum(jnp.asarray(clock_depth, dtype=jnp.int32), 1) ) parts = [left, right, sibling_edge_lr, sibling_edge_rl] g_active = ( jnp.any(merge_structural_mask) if g_structural_mask is None else jnp.asarray(g_structural_mask, dtype=bool) ) gg = ( g_projected if g_projected is not None else self.project_global(g, g_active) ) parts.append( jnp.broadcast_to(gg[None, :], left.shape[:-1] + (gg.shape[-1],)).astype( dtype ) ) if depth_feats is None: raise ValueError("tree prefix merge requires depth features") parts.append(depth_feats.astype(dtype)) clock = _route_merge_clock( level_idx, pair_idx, pair_base, self.d_model, depth_i, dtype, root_centered=True, ) parts.append(jnp.broadcast_to(clock, left.shape)) x = jnp.concatenate(parts, axis=-1) x = _tagged_rms_eqx_style( self.ln_scale, x, eps=self.ln_eps, tag_id="route.tree_prefix.merge.ln", pathway="even", kfac_structural_mask=merge_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) h = _tagged_dense( self.w1, self.b1, x, tag_id="route.tree_prefix.merge.ffn1", pathway="even", weight_eff=_we(self.w1), kfac_structural_mask=merge_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) h = fused_silu(h) delta = _tagged_dense( self.w2, self.b2, h, tag_id="route.tree_prefix.merge.ffn2", pathway="even", weight_eff=_we(self.w2), kfac_structural_mask=merge_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) raw = _tree_ngpt_residual( 0.5 * (left + right), delta, self.alpha_route, max_gain=self.ngpt_alpha_max, tag_id="route.tree_prefix.merge.alpha_route", pathway="even", kfac_structural_mask=merge_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) carry = jnp.where(left_mask[:, None] > 0, left, right) out = jnp.where(both[:, None] > 0, raw, carry) out = jnp.where( out_mask[:, None] > 0, out, jnp.zeros_like(out), ) return out, out_mask, both class _TreePrefixSelfLayer(eqx.Module): ln_scale: Float[Array, "d_model"] w_qkv: Float[Array, "d_model three_qv"] w_o: Float[Array, "d_o_in d_model"] edge_ln_scale: Float[Array, "d_model"] edge_w1: Float[Array, "d_model d_hidden"] edge_b1: Float[Array, "d_hidden"] edge_w2: Float[Array, "d_hidden h_kernel"] edge_b2: Float[Array, "h_kernel"] ffn_ln_scale: Float[Array, "d_model"] ffn_w1: Float[Array, "d_model d_ffn"] ffn_b1: Float[Array, "d_ffn"] ffn_w2: Float[Array, "d_ffn d_model"] ffn_b2: Float[Array, "d_model"] def as_tuple(self): return ( self.ln_scale, self.w_qkv, self.w_o, self.edge_ln_scale, self.edge_w1, self.edge_b1, self.edge_w2, self.edge_b2, self.ffn_ln_scale, self.ffn_w1, self.ffn_b1, self.ffn_w2, self.ffn_b2, ) class _TreePrefixCandidateLayer(eqx.Module): cand_ln_scale: Float[Array, "d_model"] prefix_ln_scale: Float[Array, "d_model"] cand_w_qv: Float[Array, "d_model two_qv"] prefix_w_kv: Float[Array, "d_model two_qv"] w_o: Float[Array, "d_o_in d_model"] edge_ln_scale: Float[Array, "d_model"] edge_w1: Float[Array, "d_model d_hidden"] edge_b1: Float[Array, "d_hidden"] edge_w2: Float[Array, "d_hidden h_kernel"] edge_b2: Float[Array, "h_kernel"] ffn_ln_scale: Float[Array, "d_model"] ffn_w1: Float[Array, "d_model d_ffn"] ffn_b1: Float[Array, "d_ffn"] ffn_w2: Float[Array, "d_ffn d_model"] ffn_b2: Float[Array, "d_model"] def as_tuple(self): return ( self.cand_ln_scale, self.prefix_ln_scale, self.cand_w_qv, self.prefix_w_kv, self.w_o, self.edge_ln_scale, self.edge_w1, self.edge_b1, self.edge_w2, self.edge_b2, self.ffn_ln_scale, self.ffn_w1, self.ffn_b1, self.ffn_w2, self.ffn_b2, ) class _HeavyRouteLayer(eqx.Module): cross_ln_scale: Float[Array, "d_model"] cross_ln_shift: Float[Array, "d_model"] cross_prefix_ln_scale: Float[Array, "d_model"] cross_prefix_ln_shift: Float[Array, "d_model"] cross_w_qv: Float[Array, "d_model two_qv"] cross_w_kv: Float[Array, "d_model two_qv"] cross_w_o: Float[Array, "d_o_in d_model"] cross_edge_ln_scale: Float[Array, "two_d_edge"] cross_edge_ln_shift: Float[Array, "two_d_edge"] cross_edge_w1: Float[Array, "two_d_edge d_heavy_edge_hidden"] cross_edge_b1: Float[Array, "d_heavy_edge_hidden"] cross_edge_w2: Float[Array, "d_heavy_edge_hidden h_kernel"] cross_edge_b2: Float[Array, "h_kernel"] self_ln_scale: Float[Array, "d_model"] self_w_qkv: Float[Array, "d_model three_qv"] self_w_o: Float[Array, "d_o_in d_model"] self_edge_ln_scale: Float[Array, "two_d_edge"] self_edge_w1: Float[Array, "two_d_edge d_heavy_edge_hidden"] self_edge_b1: Float[Array, "d_heavy_edge_hidden"] self_edge_w2: Float[Array, "d_heavy_edge_hidden h_kernel"] self_edge_b2: Float[Array, "h_kernel"] ffn_ln_scale: Float[Array, "d_model"] ffn_w1: Float[Array, "d_model d_ffn"] ffn_b1: Float[Array, "d_ffn"] ffn_w2: Float[Array, "d_ffn d_model"] ffn_b2: Float[Array, "d_model"] def as_tuple(self): return ( self.cross_ln_scale, self.cross_ln_shift, self.cross_prefix_ln_scale, self.cross_prefix_ln_shift, self.cross_w_qv, self.cross_w_kv, self.cross_w_o, self.cross_edge_ln_scale, self.cross_edge_ln_shift, self.cross_edge_w1, self.cross_edge_b1, self.cross_edge_w2, self.cross_edge_b2, self.self_ln_scale, self.self_w_qkv, self.self_w_o, self.self_edge_ln_scale, self.self_edge_w1, self.self_edge_b1, self.self_edge_w2, self.self_edge_b2, self.ffn_ln_scale, self.ffn_w1, self.ffn_b1, self.ffn_w2, self.ffn_b2, ) class _PrefixSuffixRouteBase(_RoutePointerBase): heavy_layers: list[_HeavyRouteLayer] route_prefix_suffix_layers: int = eqx.field(static=True) route_decoder_attn_impl: str = eqx.field(static=True) heavy_edge_hidden: int = eqx.field(static=True) heavy_residual_gain: float = eqx.field(static=True) def __init__( self, *, d_in: int, d_edge: int, d_global: int, d_model: int, n_heads: int, max_n: int, key: PRNGKeyArray, route_prefix_suffix_layers: int = 1, route_decoder_attn_impl: str = "mhsea_tuned", score_init_scale: float = 1.0, rope_base: float = 10000.0, rope_scaling: float = 1.0, attention_dim: int, pointer_score_dim: int, candidate_hidden: int, summary_hidden: int, ffn_hidden: int, global_tap_dim: int, ): if route_prefix_suffix_layers < 0: raise ValueError("route_prefix_suffix_layers must be >= 0") allowed_impls = {"mhsea_tuned", "einsum"} if route_decoder_attn_impl not in allowed_impls: raise ValueError( f"route_decoder_attn_impl must be one of {sorted(allowed_impls)}" ) key_base, key_heavy = jax.random.split(key) super().__init__( d_in=d_in, d_edge=d_edge, d_global=d_global, d_model=d_model, n_heads=n_heads, max_n=max_n, key=key_base, score_init_scale=score_init_scale, rope_base=rope_base, rope_scaling=rope_scaling, attention_dim=attention_dim, pointer_score_dim=pointer_score_dim, candidate_hidden=candidate_hidden, summary_hidden=summary_hidden, ffn_hidden=ffn_hidden, global_tap_dim=global_tap_dim, ) layers = int(route_prefix_suffix_layers) d_qv = self.n_heads_kernel * self.d_head d_o_in = self.n_heads * self.d_head pair_dim = 2 * self.d_edge heavy_edge_hidden = max(32, 2 * self.n_heads_kernel, 4 * self.d_edge) def w(k, shape, fan_in): return jax.random.normal(k, shape) * (fan_in**-0.5) layer_keys = jax.random.split(key_heavy, layers) heavy_layers = [] for li in range(layers): keys = jax.random.split(layer_keys[li], 11) heavy_layers.append( _HeavyRouteLayer( cross_ln_scale=jnp.ones((self.d_model,)), cross_ln_shift=jnp.zeros((self.d_model,)), cross_prefix_ln_scale=jnp.ones((self.d_model,)), cross_prefix_ln_shift=jnp.zeros((self.d_model,)), cross_w_qv=w(keys[0], (self.d_model, 2 * d_qv), self.d_model), cross_w_kv=w(keys[1], (self.d_model, 2 * d_qv), self.d_model), cross_w_o=w(keys[2], (d_o_in, self.d_model), d_o_in), cross_edge_ln_scale=jnp.ones((pair_dim,)), cross_edge_ln_shift=jnp.zeros((pair_dim,)), cross_edge_w1=w(keys[3], (pair_dim, heavy_edge_hidden), pair_dim), cross_edge_b1=jnp.zeros((heavy_edge_hidden,)), cross_edge_w2=w( keys[4], (heavy_edge_hidden, self.n_heads_kernel), heavy_edge_hidden, ), cross_edge_b2=jnp.zeros((self.n_heads_kernel,)), self_ln_scale=jnp.ones((self.d_model,)), self_w_qkv=w(keys[5], (self.d_model, 3 * d_qv), self.d_model), self_w_o=w(keys[6], (d_o_in, self.d_model), d_o_in), self_edge_ln_scale=jnp.ones((pair_dim,)), self_edge_w1=w(keys[7], (pair_dim, heavy_edge_hidden), pair_dim), self_edge_b1=jnp.zeros((heavy_edge_hidden,)), self_edge_w2=w( keys[8], (heavy_edge_hidden, self.n_heads_kernel), heavy_edge_hidden, ), self_edge_b2=jnp.zeros((self.n_heads_kernel,)), ffn_ln_scale=jnp.ones((self.d_model,)), ffn_w1=w(keys[9], (self.d_model, self.ffn_hidden), self.d_model), ffn_b1=jnp.zeros((self.ffn_hidden,)), ffn_w2=w( keys[10], (self.ffn_hidden, self.d_model), self.ffn_hidden ), ffn_b2=jnp.zeros((self.d_model,)), ) ) self.heavy_layers = heavy_layers self.route_prefix_suffix_layers = layers self.route_decoder_attn_impl = str(route_decoder_attn_impl) self.heavy_edge_hidden = heavy_edge_hidden self.heavy_residual_gain = 0.0 if layers == 0 else float(layers) ** -0.5 def _heavy_layer_params(self): return [layer.as_tuple() for layer in self.heavy_layers] def _resolve_heavy_attn_impl(self, n: int) -> str: del n return self.route_decoder_attn_impl def _heavy_edge_bias( self, edge_pair, params, *, prefix: str, structural_mask=None, scan_shared: bool = False, repeat_ndim: int | None = None, context_primal_reused_over_walkers: bool = False, ): ln_s, w1, b1, w2, b2 = params structural_mask = ( jnp.ones(edge_pair.shape[:-1], dtype=bool) if structural_mask is None else jnp.broadcast_to( jnp.asarray(structural_mask, dtype=bool), edge_pair.shape[:-1] ) ) repeat_ndim = structural_mask.ndim if repeat_ndim is None else repeat_ndim kfac_kwargs = dict( kfac_structural_mask=structural_mask, kfac_scan_shared=scan_shared, kfac_repeat_ndim=repeat_ndim, kfac_context_primal_reused_over_walkers=( context_primal_reused_over_walkers ), ) x = self._ln( ln_s, edge_pair, tag_id=f"route.heavy.{prefix}.edge_ln", **kfac_kwargs, ) x = self._dense( w1, b1, x, tag_id=f"route.heavy.{prefix}.edge_bias1", **kfac_kwargs, ) x = fused_silu(x) bias = self._dense( w2, b2, x, tag_id=f"route.heavy.{prefix}.edge_bias2", **kfac_kwargs, ) return bias def _heavy_cross_edge_bias( self, edge_pair, params, *, structural_mask=None, scan_shared: bool = False, repeat_ndim: int | None = None, context_primal_reused_over_walkers: bool = False, ): ln_s, ln_b, w1, b1, w2, b2 = params structural_mask = ( jnp.ones(edge_pair.shape[:-1], dtype=bool) if structural_mask is None else jnp.broadcast_to( jnp.asarray(structural_mask, dtype=bool), edge_pair.shape[:-1] ) ) repeat_ndim = structural_mask.ndim if repeat_ndim is None else repeat_ndim kfac_kwargs = dict( kfac_structural_mask=structural_mask, kfac_scan_shared=scan_shared, kfac_repeat_ndim=repeat_ndim, kfac_context_primal_reused_over_walkers=( context_primal_reused_over_walkers ), ) x = self._cross_ln( ln_s, ln_b, edge_pair, tag_id="route.heavy.cross.edge_ln", **kfac_kwargs, ) x = self._dense( w1, b1, x, tag_id="route.heavy.cross.edge_bias1", **kfac_kwargs, ) x = fused_silu(x) return self._dense( w2, b2, x, tag_id="route.heavy.cross.edge_bias2", **kfac_kwargs, ) def _route_attention( self, q: Float[Array, "b n h d_head"], k: Float[Array, "b n h d_head"], v: Float[Array, "b n h d_head"], edge_bias: Float[Array, "b n n h"], key_mask: Int[Array, "b n"] | Array, *, impl: str, key_mask_only: bool = False, attention_mask: Array | None = None, sequence_axis_name=None, sequence_mesh=None, ) -> Float[Array, "b n h d_head"]: dtype = q.dtype valid = key_mask.astype(bool)[:, None, :] if attention_mask is not None: valid = valid & attention_mask.astype(bool) has_key = jnp.any(valid, axis=-1) if impl == "einsum": q_c = q k_c = k v_c = v bias_c = edge_bias if sequence_axis_name is not None: from jax.sharding import NamedSharding, PartitionSpec as P def _sharding(*axes): spec = P(*axes) return ( NamedSharding(sequence_mesh, spec) if sequence_mesh is not None else spec ) q_c = jax.lax.with_sharding_constraint( q_c, _sharding(None, sequence_axis_name, None, None), ) k_c = jax.lax.with_sharding_constraint( k_c, _sharding(None, None, None, None), ) v_c = jax.lax.with_sharding_constraint( v_c, _sharding(None, None, None, None), ) bias_c = jax.lax.with_sharding_constraint( bias_c, _sharding(None, sequence_axis_name, None, None), ) logits = jnp.einsum("bihd,bjhd->bhij", q_c, k_c) logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=dtype)) logits = logits + jnp.transpose(bias_c, (0, 3, 1, 2)) if sequence_axis_name is not None: logits = jax.lax.with_sharding_constraint( logits, _sharding(None, None, sequence_axis_name, None), ) logits = jnp.where( valid[:, None, :, :], logits, jnp.asarray(-1.0e30, dtype=dtype), ) if sequence_axis_name is not None: logits = jax.lax.with_sharding_constraint( logits, _sharding(None, None, sequence_axis_name, None), ) alpha = jax.nn.softmax(logits, axis=-1) if sequence_axis_name is not None: alpha = jax.lax.with_sharding_constraint( alpha, _sharding(None, None, sequence_axis_name, None), ) out = jnp.einsum("bhij,bjhd->bihd", alpha, v_c) if sequence_axis_name is not None: out = jax.lax.with_sharding_constraint( out, _sharding(None, sequence_axis_name, None, None), ) elif impl == "mhsea_tuned": from hamiltonzero.model.pallas_attention import mhsea_tuned_edge_attention if key_mask_only or attention_mask is not None: edge_bias = jnp.where( valid[..., None], edge_bias, jnp.asarray(-1.0e30, dtype=edge_bias.dtype), ) key_mask = jnp.ones_like(key_mask) d_head_padded = max(16, self.d_head) pad_amount = d_head_padded - self.d_head if pad_amount: scale = jnp.sqrt(jnp.asarray(d_head_padded / self.d_head, dtype=dtype)) q = jnp.concatenate( [ q * scale, jnp.zeros(q.shape[:-1] + (pad_amount,), dtype=q.dtype), ], axis=-1, ) k = jnp.concatenate( [ k, jnp.zeros(k.shape[:-1] + (pad_amount,), dtype=k.dtype), ], axis=-1, ) v = jnp.concatenate( [ v, jnp.zeros(v.shape[:-1] + (pad_amount,), dtype=v.dtype), ], axis=-1, ) out = jax.vmap( lambda q_b, k_b, v_b, bias_b, mask_b: mhsea_tuned_edge_attention( q_b, k_b, v_b, bias_b, mask_b.astype(jnp.int32) ) )(q, k, v, edge_bias, key_mask) out = out[..., : self.d_head] else: raise ValueError("route attention must be 'einsum' or 'mhsea_tuned'") return jnp.where(has_key[..., None, None], out, jnp.zeros_like(out)) def _collapse_heavy_heads(self, out): gate_heads = out[..., : self.n_heads, :] value_heads = out[..., self.n_heads :, :] out = jax.nn.sigmoid(gate_heads) * value_heads return out.reshape(out.shape[:-2] + (self.n_heads * self.d_head,)) def _heavy_prefix_pairs_teacher(self, edge, perm): n = edge.shape[0] edge_i_p = edge[:, perm, :] edge_p_i = jnp.transpose(edge[perm, :, :], (1, 0, 2)) pair = jnp.concatenate([edge_i_p, edge_p_i], axis=-1) return pair def _heavy_suffix_pairs_teacher(self, edge): n = edge.shape[0] idx = jnp.arange(n, dtype=jnp.int32) edge_i_j = edge[idx[:, None], idx[None, :], :] edge_j_i = edge[idx[None, :], idx[:, None], :] pair = jnp.concatenate([edge_i_j, edge_j_i], axis=-1) return pair def _heavy_layer_teacher( self, cand: Float[Array, "n n d_model"], z: Float[Array, "n d_model"], edge: Float[Array, "n n d_edge"], perm: Int[Array, "n"], mask: Int[Array, "n"] | Array, pos_of_node: Int[Array, "n"], params, *, impl: str, ) -> Float[Array, "n n d_model"]: ( cross_ln_s, cross_ln_b, cross_prefix_ln_s, cross_prefix_ln_b, cross_w_qv, cross_w_kv, cross_w_o, cross_edge_ln_s, cross_edge_ln_b, cross_edge_w1, cross_edge_b1, cross_edge_w2, cross_edge_b2, self_ln_s, self_w_qkv, self_w_o, self_edge_ln_s, self_edge_w1, self_edge_b1, self_edge_w2, self_edge_b2, ffn_ln_s, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ) = params n = cand.shape[0] dtype = cand.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) candidate_structural_mask = ( mask_bool[:, None] & mask_bool[None, :] & (pos_of_node[None, :] >= idx[:, None]) ) prefix_structural_mask = mask_bool prefix_pair_structural_mask = ( mask_bool[:, None] & mask_bool[None, :] & (idx[None, :] < pos_of_node[:, None]) ) suffix_pair_structural_mask = mask_bool[:, None] & mask_bool[None, :] context_reuse = True candidate_kfac = dict( kfac_structural_mask=candidate_structural_mask, kfac_repeat_ndim=2, kfac_context_primal_reused_over_walkers=context_reuse, ) prefix_kfac = dict( kfac_structural_mask=prefix_structural_mask, kfac_repeat_ndim=1, kfac_context_primal_reused_over_walkers=context_reuse, ) query_mask = mask.astype(dtype)[None, :, None] x_ln = self._cross_ln( cross_ln_s, cross_ln_b, cand, tag_id="route.heavy.cross.ln", **candidate_kfac, ) qv = self._dense_no_bias( cross_w_qv, x_ln, tag_id="route.heavy.cross.qv", **candidate_kfac, ).reshape(n, n, 2, self.n_heads_kernel, self.d_head) q = qv[:, :, 0] v_self = qv[:, :, 1] z_ln = self._cross_ln( cross_prefix_ln_s, cross_prefix_ln_b, z, tag_id="route.heavy.cross.prefix_ln", **prefix_kfac, ) kv = self._dense_no_bias( cross_w_kv, z_ln, tag_id="route.heavy.cross.kv", **prefix_kfac, ).reshape(n, 2, self.n_heads_kernel, self.d_head) k = kv[:, 0] v = kv[:, 1] k_b = jnp.broadcast_to(k[None, :, :, :], q.shape) v_b = jnp.broadcast_to(v[None, :, :, :], q.shape) prefix_pairs = self._heavy_prefix_pairs_teacher(edge, perm) cross_bias = self._heavy_cross_edge_bias( prefix_pairs, ( cross_edge_ln_s, cross_edge_ln_b, cross_edge_w1, cross_edge_b1, cross_edge_w2, cross_edge_b2, ), structural_mask=prefix_pair_structural_mask, repeat_ndim=2, context_primal_reused_over_walkers=context_reuse, ) lca_tk = lca_alibi_bias( idx, idx, lca_fixed_slopes(self.n_heads_kernel, dtype=cand.dtype), ) pos_bias = jnp.transpose(lca_tk, (1, 2, 0)) cross_bias = cross_bias[None, :, :, :] + pos_bias[:, None, :, :] key_mask = mask.astype(bool)[None, :] & (idx[None, :] < idx[:, None]) cross_out = self._route_attention( q, k_b, v_b, cross_bias, key_mask, impl=impl, key_mask_only=True, ) cross_flat = self._collapse_heavy_heads(cross_out) delta = self._dense_no_bias( cross_w_o, cross_flat, tag_id="route.heavy.cross.o", **candidate_kfac, ) cand = cand + query_mask * self.heavy_residual_gain * delta x_ln = self._ln( self_ln_s, cand, tag_id="route.heavy.self.ln", **candidate_kfac, ) qkv = self._dense_no_bias( self_w_qkv, x_ln, tag_id="route.heavy.self.qkv", **candidate_kfac, ).reshape(n, n, 3, self.n_heads_kernel, self.d_head) q = qkv[:, :, 0] k = qkv[:, :, 1] v = qkv[:, :, 2] suffix_pairs = self._heavy_suffix_pairs_teacher(edge) suffix_bias = self._heavy_edge_bias( suffix_pairs, ( self_edge_ln_s, self_edge_w1, self_edge_b1, self_edge_w2, self_edge_b2, ), prefix="self", structural_mask=suffix_pair_structural_mask, repeat_ndim=2, context_primal_reused_over_walkers=context_reuse, ) suffix_bias = jnp.broadcast_to( suffix_bias[None, :, :, :], (n,) + suffix_bias.shape, ) suffix_mask = mask.astype(bool)[None, :] & ( pos_of_node[None, :] >= idx[:, None] ) self_out = self._route_attention( q, k, v, suffix_bias, suffix_mask, impl=impl, ) self_flat = self._collapse_heavy_heads(self_out) delta = self._dense_no_bias( self_w_o, self_flat, tag_id="route.heavy.self.o", **candidate_kfac, ) cand = cand + query_mask * self.heavy_residual_gain * delta ffn_in = self._ln( ffn_ln_s, cand, tag_id="route.heavy.ffn.ln", **candidate_kfac, ) ffn = self._dense( ffn_w1, ffn_b1, ffn_in, tag_id="route.heavy.ffn1", **candidate_kfac, ) ffn = fused_silu(ffn) delta = self._dense( ffn_w2, ffn_b2, ffn, tag_id="route.heavy.ffn2", **candidate_kfac, ) return cand + query_mask * self.heavy_residual_gain * delta def _apply_heavy_teacher(self, base, z, edge, perm, mask): n = base.shape[0] idx = jnp.arange(n, dtype=jnp.int32) pos_of_node = jnp.zeros((n,), dtype=jnp.int32).at[perm].set(idx) impl = self._resolve_heavy_attn_impl(n) def apply_one(cand, params): return self._heavy_layer_teacher( cand, z, edge, perm, mask, pos_of_node, params, impl=impl, ) params = self._heavy_layer_params() cand = base for layer in params: cand = apply_one(cand, layer) return cand def _heavy_prefix_pairs_step( self, edge, prefix_ids, *, edge_transpose=None, ): edge_i_p = edge[:, prefix_ids, :] edge_p_i = ( jnp.transpose(edge[prefix_ids, :, :], (1, 0, 2)) if edge_transpose is None else edge_transpose[:, prefix_ids, :] ) return jnp.concatenate([edge_i_p, edge_p_i], axis=-1) def _heavy_suffix_pairs_step(self, edge, *, edge_transpose=None): edge_i_j = edge edge_j_i = ( jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose ) return jnp.concatenate([edge_i_j, edge_j_i], axis=-1) def _heavy_cross_biases(self, edge): all_pairs = self._heavy_suffix_pairs_step(edge) params = self._heavy_layer_params() return tuple( self._heavy_cross_edge_bias( all_pairs, layer[7:13], context_primal_reused_over_walkers=True, ) for layer in params ) def _heavy_suffix_biases(self, edge): suffix_pairs = self._heavy_suffix_pairs_step(edge) params = self._heavy_layer_params() return tuple( self._heavy_edge_bias( suffix_pairs, layer[16:21], prefix="self", context_primal_reused_over_walkers=True, ) for layer in params ) def _heavy_biases_tiled( self, edge, edge_transpose, *, pair_tile_size: int, sequence_axis_name: str | None = None, sequence_mesh=None, ): n = edge.shape[0] tile = min(int(pair_tile_size), n) if tile < 1: raise ValueError("pair_tile_size must be positive") n_tiles = (n + tile - 1) // tile padded_n = n_tiles * tile source_pad = padded_n - n edge_transpose = ( jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose ) edge_padded = jnp.pad(edge, ((0, 0), (0, source_pad), (0, 0))) edge_transpose_padded = jnp.pad( edge_transpose, ((0, 0), (0, source_pad), (0, 0)), ) def _seq_constraint(value, *axes): if sequence_axis_name is None: return value from jax.sharding import NamedSharding, PartitionSpec as P spec = P(*axes) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) params = self._heavy_layer_params() layer_params = tuple(params) def project_layer(layer, param_slice, project_bias): output0 = _seq_constraint( jnp.zeros( (n, padded_n, self.n_heads_kernel), dtype=edge.dtype, ), sequence_axis_name, None, None, ) def project_tile(tile_index, output): start = tile_index * tile edge_tile = jax.lax.dynamic_slice_in_dim( edge_padded, start, tile, axis=1, ) edge_transpose_tile = jax.lax.dynamic_slice_in_dim( edge_transpose_padded, start, tile, axis=1, ) edge_pair = jnp.concatenate( [edge_tile, edge_transpose_tile], axis=-1, ) edge_pair = _seq_constraint( edge_pair, sequence_axis_name, None, None, ) bias = project_bias( edge_pair, layer[param_slice], context_primal_reused_over_walkers=True, ) output = jax.lax.dynamic_update_slice_in_dim( output, bias, start, axis=1, ) return _seq_constraint( output, sequence_axis_name, None, None, ) output = jax.lax.fori_loop( 0, n_tiles, project_tile, output0, ) return output[:, :n, :] cross_biases = tuple( project_layer(layer, slice(7, 13), self._heavy_cross_edge_bias) for layer in layer_params ) suffix_biases = tuple( project_layer( layer, slice(16, 21), lambda edge_pair, params, **kwargs: self._heavy_edge_bias( edge_pair, params, prefix="self", **kwargs ), ) for layer in layer_params ) return cross_biases, suffix_biases def _pack_heavy_static_bias_tables(self, edge): return self._heavy_cross_biases(edge) + self._heavy_suffix_biases(edge) def _unpack_heavy_static_bias_tables(self, tables): layers = int(self.route_prefix_suffix_layers) expected = 2 * layers if len(tables) != expected: raise ValueError( "RouterStatic static_bias_tables has " f"{len(tables)} leaves; expected {expected} for " f"route_prefix_suffix_layers={layers}" ) return tables[:layers], tables[layers:] def _heavy_layer_step_candidates( self, cand: Float[Array, "n d_model"], hidden_cache: Float[Array, "n d_model"], edge: Float[Array, "n n d_edge"], prefix_ids: Int[Array, "n"], picked: Array, mask: Int[Array, "n"] | Array, t: Int[Array, ""], params, *, impl: str, cross_bias=None, suffix_bias=None, edge_transpose=None, sequence_axis_name=None, sequence_mesh=None, ) -> Float[Array, "n d_model"]: ( cross_ln_s, cross_ln_b, cross_prefix_ln_s, cross_prefix_ln_b, cross_w_qv, cross_w_kv, cross_w_o, cross_edge_ln_s, cross_edge_ln_b, cross_edge_w1, cross_edge_b1, cross_edge_w2, cross_edge_b2, self_ln_s, self_w_qkv, self_w_o, self_edge_ln_s, self_edge_w1, self_edge_b1, self_edge_w2, self_edge_b2, ffn_ln_s, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ) = params n = cand.shape[0] dtype = cand.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) row_active = mask_bool[t] candidate_structural_mask = row_active & mask_bool & (~picked.astype(bool)) prefix_structural_mask = row_active & mask_bool & (idx < t) cross_pair_structural_mask = ( candidate_structural_mask[:, None] & prefix_structural_mask[None, :] ) self_pair_structural_mask = ( candidate_structural_mask[:, None] & candidate_structural_mask[None, :] ) context_reuse = True candidate_kfac = dict( kfac_structural_mask=candidate_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, kfac_context_primal_reused_over_walkers=context_reuse, ) prefix_kfac = dict( kfac_structural_mask=prefix_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, kfac_context_primal_reused_over_walkers=context_reuse, ) query_mask = mask.astype(dtype).reshape(n, 1) def _seq_constraint(value, *axes): if sequence_axis_name is None: return value from jax.sharding import NamedSharding, PartitionSpec as P spec = P(*axes) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) cand = _seq_constraint(cand, sequence_axis_name, None) x_ln = self._cross_ln( cross_ln_s, cross_ln_b, cand, tag_id="route.heavy.cross.ln", **candidate_kfac, ) qv = self._dense_no_bias( cross_w_qv, x_ln, tag_id="route.heavy.cross.qv", **candidate_kfac, ).reshape(n, 2, self.n_heads_kernel, self.d_head) qv = _seq_constraint( qv, sequence_axis_name, None, None, None, ) q = qv[:, 0] v_self = qv[:, 1] z_ln = self._cross_ln( cross_prefix_ln_s, cross_prefix_ln_b, hidden_cache, tag_id="route.heavy.cross.prefix_ln", **prefix_kfac, ) kv = self._dense_no_bias( cross_w_kv, z_ln, tag_id="route.heavy.cross.kv", **prefix_kfac, ).reshape(n, 2, self.n_heads_kernel, self.d_head) kv = _seq_constraint(kv, None, None, None, None) k = kv[:, 0] v = kv[:, 1] q = _seq_constraint(q, sequence_axis_name, None, None) v_self = _seq_constraint( v_self, sequence_axis_name, None, None, ) k = _seq_constraint(k, None, None, None) v = _seq_constraint(v, None, None, None) if cross_bias is None: prefix_pairs = self._heavy_prefix_pairs_step( edge, prefix_ids, edge_transpose=edge_transpose, ) cross_bias = self._heavy_cross_edge_bias( prefix_pairs, ( cross_edge_ln_s, cross_edge_ln_b, cross_edge_w1, cross_edge_b1, cross_edge_w2, cross_edge_b2, ), structural_mask=cross_pair_structural_mask, scan_shared=True, repeat_ndim=2, context_primal_reused_over_walkers=context_reuse, ) else: cross_bias = cross_bias[:, prefix_ids, :] cross_bias = _seq_constraint( cross_bias, sequence_axis_name, None, None, ) pos_bias = lca_alibi_bias( jnp.asarray([t], jnp.int32), idx, lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), )[:, 0, :] cross_bias = cross_bias + jnp.transpose(pos_bias, (1, 0))[None, :, :] key_mask = mask.astype(bool) & (idx < t) cross_out = self._route_attention( q[None, :, :, :], k[None, :, :, :], v[None, :, :, :], cross_bias[None, :, :, :], key_mask[None, :], impl=impl, key_mask_only=True, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, )[0] cross_out = _seq_constraint( cross_out, sequence_axis_name, None, None, ) cross_flat = self._collapse_heavy_heads(cross_out) delta = self._dense_no_bias( cross_w_o, cross_flat, tag_id="route.heavy.cross.o", **candidate_kfac, ) cand = cand + query_mask * self.heavy_residual_gain * delta x_ln = self._ln( self_ln_s, cand, tag_id="route.heavy.self.ln", **candidate_kfac, ) qkv = self._dense_no_bias( self_w_qkv, x_ln, tag_id="route.heavy.self.qkv", **candidate_kfac, ).reshape(n, 3, self.n_heads_kernel, self.d_head) qkv = _seq_constraint( qkv, sequence_axis_name, None, None, None, ) q = qkv[:, 0] k = qkv[:, 1] v = qkv[:, 2] q = _seq_constraint(q, sequence_axis_name, None, None) k = _seq_constraint(k, None, None, None) v = _seq_constraint(v, None, None, None) if suffix_bias is None: suffix_pairs = self._heavy_suffix_pairs_step( edge, edge_transpose=edge_transpose, ) suffix_bias = self._heavy_edge_bias( suffix_pairs, ( self_edge_ln_s, self_edge_w1, self_edge_b1, self_edge_w2, self_edge_b2, ), prefix="self", structural_mask=self_pair_structural_mask, scan_shared=True, repeat_ndim=2, context_primal_reused_over_walkers=context_reuse, ) suffix_bias = _seq_constraint( suffix_bias, sequence_axis_name, None, None, ) suffix_mask = mask.astype(bool) & (~picked) self_out = self._route_attention( q[None, :, :, :], k[None, :, :, :], v[None, :, :, :], suffix_bias[None, :, :, :], suffix_mask[None, :], impl=impl, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, )[0] self_out = _seq_constraint( self_out, sequence_axis_name, None, None, ) self_flat = self._collapse_heavy_heads(self_out) delta = self._dense_no_bias( self_w_o, self_flat, tag_id="route.heavy.self.o", **candidate_kfac, ) cand = cand + query_mask * self.heavy_residual_gain * delta ffn_in = self._ln( ffn_ln_s, cand, tag_id="route.heavy.ffn.ln", **candidate_kfac, ) ffn = self._dense( ffn_w1, ffn_b1, ffn_in, tag_id="route.heavy.ffn1", **candidate_kfac, ) ffn = fused_silu(ffn) delta = self._dense( ffn_w2, ffn_b2, ffn, tag_id="route.heavy.ffn2", **candidate_kfac, ) return cand + query_mask * self.heavy_residual_gain * delta def _apply_heavy_step( self, base, hidden_cache, edge, prefix_ids, picked, mask, t, *, cross_biases=None, suffix_biases=None, edge_transpose=None, sequence_axis_name=None, sequence_mesh=None, ): impl = self._resolve_heavy_attn_impl(base.shape[0]) layer_params = self._heavy_layer_params() n_layers = int(self.route_prefix_suffix_layers) if cross_biases is None or len(cross_biases) == 0: cross_biases = (None,) * n_layers if suffix_biases is None: suffix_biases = (None,) * n_layers cand = base for params, cross_bias, suffix_bias in zip( layer_params, cross_biases, suffix_biases, ): cand = self._heavy_layer_step_candidates( cand, hidden_cache, edge, prefix_ids, picked, mask, t, params, impl=impl, cross_bias=cross_bias, suffix_bias=suffix_bias, edge_transpose=edge_transpose, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) return cand class TreePrefixPointerMHSEA(_PrefixSuffixRouteBase): tree_merge: _TreePrefixMerge tree_edge_merge: EdgeMergeOp tree_level_layer: _TreePrefixSelfLayer tree_level_fwl: CausalRouterEdgeFWLUpdate alpha_tree_level_attn: Float[Array, "d_model"] alpha_tree_level_ffn: Float[Array, "d_model"] tree_prefix_layers: list[_TreePrefixSelfLayer] tree_candidate_layers: list[_TreePrefixCandidateLayer] g_step_pool: "GDescriptorPool" g_step_update: "TreeGlobalUpdate" g_prefix_ffn_w: Float[Array, "d_gstream d_model"] g_cand_ffn_w: Float[Array, "d_gstream d_model"] tree_edge_ln_scale: Float[Array, "two_d_edge"] tree_edge_w1: Float[Array, "two_d_edge d_msg_hidden"] tree_edge_b1: Float[Array, "d_msg_hidden"] tree_edge_w2: Float[Array, "d_msg_hidden d_model"] tree_edge_b2: Float[Array, "d_model"] route_tree_prefix_layers: int = eqx.field(static=True) route_tree_prefix_candidate_layers: int = eqx.field(static=True) tree_prefix_merge_hidden: int = eqx.field(static=True) tree_prefix_edge_hidden: int = eqx.field(static=True) tree_prefix_residual_gain: float = eqx.field(static=True) tree_candidate_residual_gain: float = eqx.field(static=True) tree_ngpt_alpha_max: float = eqx.field(static=True) def __init__( self, *, d_in: int, d_edge: int, d_global: int, d_model: int, n_heads: int, max_n: int, key: PRNGKeyArray, route_tree_prefix_layers: int = 1, route_tree_prefix_candidate_layers: int = 1, route_tree_prefix_merge_hidden: int, route_tree_prefix_post_prefix_suffix_layers: int = 0, score_init_scale: float = 1.0, route_decoder_attn_impl: str = "mhsea_tuned", rope_base: float = 10000.0, rope_scaling: float = 1.0, attention_dim: int, pointer_score_dim: int, candidate_hidden: int, summary_hidden: int, ffn_hidden: int, global_tap_dim: int, alpha_init: float, alpha_max: float, ): if route_tree_prefix_layers < 0: raise ValueError("route_tree_prefix_layers must be >= 0") if route_tree_prefix_candidate_layers < 0: raise ValueError("route_tree_prefix_candidate_layers must be >= 0") if route_tree_prefix_post_prefix_suffix_layers < 0: raise ValueError("route_tree_prefix_post_prefix_suffix_layers must be >= 0") key_base, key_tree = jax.random.split(key) super().__init__( d_in=d_in, d_edge=d_edge, d_global=d_global, d_model=d_model, n_heads=n_heads, max_n=max_n, key=key_base, score_init_scale=score_init_scale, route_prefix_suffix_layers=int(route_tree_prefix_post_prefix_suffix_layers), route_decoder_attn_impl=route_decoder_attn_impl, rope_base=rope_base, rope_scaling=rope_scaling, attention_dim=attention_dim, pointer_score_dim=pointer_score_dim, candidate_hidden=candidate_hidden, summary_hidden=summary_hidden, ffn_hidden=ffn_hidden, global_tap_dim=global_tap_dim, ) layers = int(route_tree_prefix_layers) cand_layers = int(route_tree_prefix_candidate_layers) merge_hidden = int(route_tree_prefix_merge_hidden) d_qv = self.n_heads_kernel * self.d_head d_o_in = self.n_heads * self.d_head edge_hidden = max(32, 2 * self.n_heads_kernel, self.msg_hidden) msg_hidden = self.msg_hidden k_edge, k_merge, k_level, k_prefix, k_cand = jax.random.split(key_tree, 5) def w(k, shape, fan_in): return jax.random.normal(k, shape) * (fan_in**-0.5) self.tree_edge_ln_scale = jnp.ones((2 * self.d_edge,)) ek1, ek2 = jax.random.split(k_edge) self.tree_edge_w1 = w(ek1, (2 * self.d_edge, msg_hidden), 2 * self.d_edge) self.tree_edge_b1 = jnp.zeros((msg_hidden,)) self.tree_edge_w2 = w(ek2, (msg_hidden, self.d_model), msg_hidden) self.tree_edge_b2 = jnp.zeros((self.d_model,)) self.tree_merge = _TreePrefixMerge( self.d_model, hidden=merge_hidden, max_depth=max(1, default_tree_depth(max_n)), key=k_merge, ln_eps=1.0e-5, gladder_d_g=int(self.d_global), alpha_init=alpha_init, alpha_max=alpha_max, ) self.tree_edge_merge = EdgeMergeOp( d_edge=self.d_model, d_c=self.d_model, key=jax.random.fold_in(k_level, 0xE06E), alpha_init=alpha_init, alpha_max=alpha_max, d_hidden=None, n_blocks=2, edge_node_ctx_dim=None, ) from .global_ladder import GDescriptorPool, TreeGlobalUpdate _k_rg = jax.random.split(jax.random.fold_in(k_merge, 0x61B6), 4) self.g_step_pool = GDescriptorPool( int(self.d_global), self.d_model, key=_k_rg[0], tag="gladder.route.step.pool", ) self.g_step_update = TreeGlobalUpdate( int(self.d_global), self.g_step_pool.d_out, key=_k_rg[1], tag="gladder.route.step.upd", tap_dim=global_tap_dim, alpha_init=alpha_init, alpha_max=alpha_max, ) self.g_prefix_ffn_w = jax.random.normal( _k_rg[2], (int(self.d_global), self.d_model) ) * (int(self.d_global) ** -0.5) self.g_cand_ffn_w = jax.random.normal( _k_rg[3], (int(self.d_global), self.d_model) ) * (int(self.d_global) ** -0.5) ks = jax.random.split(k_level, 7) self.tree_level_layer = _TreePrefixSelfLayer( ln_scale=jnp.ones((self.d_model,)), w_qkv=w(ks[0], (self.d_model, 3 * d_qv), self.d_model), w_o=w(ks[1], (d_o_in, self.d_model), d_o_in), edge_ln_scale=jnp.ones((self.d_model,)), edge_w1=w(ks[2], (self.d_model, edge_hidden), self.d_model), edge_b1=jnp.zeros((edge_hidden,)), edge_w2=w(ks[3], (edge_hidden, self.n_heads_kernel), edge_hidden), edge_b2=jnp.zeros((self.n_heads_kernel,)), ffn_ln_scale=jnp.ones((self.d_model,)), ffn_w1=w(ks[4], (self.d_model, self.ffn_hidden), self.d_model), ffn_b1=jnp.zeros((self.ffn_hidden,)), ffn_w2=w(ks[5], (self.ffn_hidden, self.d_model), self.ffn_hidden), ffn_b2=jnp.zeros((self.d_model,)), ) self.tree_level_fwl = CausalRouterEdgeFWLUpdate( d_c=self.d_model, d_edge=self.d_model, channels=max(32, self.d_model // 2), alpha_init=alpha_init, alpha_max=alpha_max, key=ks[6], ) self.alpha_tree_level_attn = float(alpha_init) * jnp.ones((self.d_model,)) self.alpha_tree_level_ffn = float(alpha_init) * jnp.ones((self.d_model,)) self.tree_ngpt_alpha_max = float(alpha_max) prefix_keys = jax.random.split(k_prefix, max(layers, 1)) prefix_layers = [] for li in range(layers): ks = jax.random.split(prefix_keys[li], 7) prefix_layers.append( _TreePrefixSelfLayer( ln_scale=jnp.ones((self.d_model,)), w_qkv=w(ks[0], (self.d_model, 3 * d_qv), self.d_model), w_o=w(ks[1], (d_o_in, self.d_model), d_o_in), edge_ln_scale=jnp.ones((self.d_model,)), edge_w1=w(ks[2], (self.d_model, edge_hidden), self.d_model), edge_b1=jnp.zeros((edge_hidden,)), edge_w2=w(ks[3], (edge_hidden, self.n_heads_kernel), edge_hidden), edge_b2=jnp.zeros((self.n_heads_kernel,)), ffn_ln_scale=jnp.ones((self.d_model,)), ffn_w1=w(ks[4], (self.d_model, self.ffn_hidden), self.d_model), ffn_b1=jnp.zeros((self.ffn_hidden,)), ffn_w2=w(ks[5], (self.ffn_hidden, self.d_model), self.ffn_hidden), ffn_b2=jnp.zeros((self.d_model,)), ) ) self.tree_prefix_layers = prefix_layers cand_keys = jax.random.split(k_cand, max(cand_layers, 1)) tree_candidate_layers = [] for li in range(cand_layers): ks = jax.random.split(cand_keys[li], 8) tree_candidate_layers.append( _TreePrefixCandidateLayer( cand_ln_scale=jnp.ones((self.d_model,)), prefix_ln_scale=jnp.ones((self.d_model,)), cand_w_qv=w(ks[0], (self.d_model, 2 * d_qv), self.d_model), prefix_w_kv=w(ks[1], (self.d_model, 2 * d_qv), self.d_model), w_o=w(ks[2], (d_o_in, self.d_model), d_o_in), edge_ln_scale=jnp.ones((self.d_model,)), edge_w1=w(ks[3], (self.d_model, edge_hidden), self.d_model), edge_b1=jnp.zeros((edge_hidden,)), edge_w2=w(ks[4], (edge_hidden, self.n_heads_kernel), edge_hidden), edge_b2=jnp.zeros((self.n_heads_kernel,)), ffn_ln_scale=jnp.ones((self.d_model,)), ffn_w1=w(ks[5], (self.d_model, self.ffn_hidden), self.d_model), ffn_b1=jnp.zeros((self.ffn_hidden,)), ffn_w2=w(ks[6], (self.ffn_hidden, self.d_model), self.ffn_hidden), ffn_b2=jnp.zeros((self.d_model,)), ) ) self.tree_candidate_layers = tree_candidate_layers self.route_tree_prefix_layers = layers self.route_tree_prefix_candidate_layers = cand_layers self.tree_prefix_merge_hidden = int(merge_hidden) self.tree_prefix_edge_hidden = int(edge_hidden) self.tree_prefix_residual_gain = 0.0 if layers == 0 else float(layers) ** -0.5 self.tree_candidate_residual_gain = ( 0.0 if cand_layers == 0 else float(cand_layers) ** -0.5 ) def _tree_prefix_layer_params(self): return [layer.as_tuple() for layer in self.tree_prefix_layers] def _tree_candidate_layer_params(self): return [layer.as_tuple() for layer in self.tree_candidate_layers] def _resolve_tree_attn_impl(self) -> str: return self.route_decoder_attn_impl def _tree_edge_message_mlp(self, edge_pair, structural_mask): x = self._ln( self.tree_edge_ln_scale, edge_pair, tag_id="route.tree_prefix.edge_msg_ln", kfac_structural_mask=structural_mask, kfac_repeat_ndim=2, ) x = self._dense( self.tree_edge_w1, self.tree_edge_b1, x, tag_id="route.tree_prefix.edge_msg1", kfac_structural_mask=structural_mask, kfac_repeat_ndim=2, ) x = fused_silu(x) return self._dense( self.tree_edge_w2, self.tree_edge_b2, x, tag_id="route.tree_prefix.edge_msg2", kfac_structural_mask=structural_mask, kfac_repeat_ndim=2, ) def _tree_pair_messages(self, edge, mask): edge_pair = jnp.concatenate([jnp.swapaxes(edge, 0, 1), edge], axis=-1) mask_bool = mask.astype(bool) structural_mask = mask_bool[:, None] & mask_bool[None, :] return self._tree_edge_message_mlp(edge_pair, structural_mask) def _tree_pair_messages_for_route( self, edge, route_ids, mask, *, edge_transpose=None, sequence_axis_name=None, sequence_mesh=None, row_permute_fn=None, ): edge_transpose = ( jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose ) edge_rows = ( jnp.take(edge, route_ids, axis=0) if row_permute_fn is None else row_permute_fn(edge, route_ids) ) edge_transpose_rows = ( jnp.take(edge_transpose, route_ids, axis=0) if row_permute_fn is None else row_permute_fn(edge_transpose, route_ids) ) edge_fwd = jnp.take(edge_rows, route_ids, axis=1) edge_rev = jnp.take(edge_transpose_rows, route_ids, axis=1) if sequence_axis_name is not None: from jax.sharding import NamedSharding, PartitionSpec as P spec = P(sequence_axis_name, None, None) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) edge_fwd = jax.lax.with_sharding_constraint(edge_fwd, spec) edge_rev = jax.lax.with_sharding_constraint(edge_rev, spec) edge_pair = jnp.concatenate([edge_rev, edge_fwd], axis=-1) route_mask = mask.astype(bool)[route_ids] structural_mask = route_mask[:, None] & route_mask[None, :] return self._tree_edge_message_mlp(edge_pair, structural_mask) def _tree_pair_message_row( self, edge, source, mask, *, edge_transpose=None, ): edge_transpose = ( jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose ) edge_pair = jnp.concatenate( [edge[:, source, :], edge_transpose[:, source, :]], axis=-1, ) structural_mask = mask.astype(bool)[source] & mask.astype(bool) return self._tree_edge_message_mlp(edge_pair, structural_mask) def _tree_pair_message_column( self, edge, destination, mask, *, edge_transpose=None, ): edge_transpose = ( jnp.swapaxes(edge, 0, 1) if edge_transpose is None else edge_transpose ) edge_pair = jnp.concatenate( [edge_transpose[:, destination, :], edge[:, destination, :]], axis=-1, ) structural_mask = mask.astype(bool)[destination] & mask.astype(bool) return self._tree_edge_message_mlp(edge_pair, structural_mask) def _tree_clock_depth_from_mask(self, mask): n_active = jnp.maximum(jnp.sum(mask.astype(jnp.int32)), 1) depth = jnp.ceil(jnp.log2(n_active.astype(jnp.float32))).astype(jnp.int32) return jnp.maximum(depth, 1) def _apply_tree_level_attention(self, nodes, edges, mask, level_idx): edges = self.tree_level_fwl.apply_residual( edges, nodes, mask, kfac_scan_shared=True, ) ( ln_s, w_qkv, w_o, edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2, ffn_ln_s, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ) = self.tree_level_layer.as_tuple() del level_idx n = nodes.shape[0] dtype = nodes.dtype mask_bool = mask.astype(bool) idx = jnp.arange(n, dtype=jnp.int32) node_structural_mask = mask_bool pair_structural_mask = ( mask_bool[:, None] & mask_bool[None, :] & (idx[None, :] <= idx[:, None]) ) query_mask = mask.astype(dtype)[:, None] x_ln = self._ln( ln_s, nodes, tag_id="route.tree_prefix.level.ln", kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) qkv = self._dense_no_bias( w_qkv, x_ln, tag_id="route.tree_prefix.level.qkv", kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ).reshape(n, 3, self.n_heads_kernel, self.d_head) q = qkv[:, 0] k = qkv[:, 1] v = qkv[:, 2] edge_bias = self._tree_prefix_edge_bias( edges, (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), prefix="level", kfac_structural_mask=pair_structural_mask, kfac_scan_shared=True, ) edge_bias = edge_bias + jnp.transpose( lca_alibi_bias( idx, idx, lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), ), (1, 2, 0), ) compute_dtype = dtype q_c = q.astype(compute_dtype) k_c = k.astype(compute_dtype) v_c = v.astype(compute_dtype) logits = jnp.einsum("ihd,jhd->hij", q_c, k_c) logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=compute_dtype)) logits = logits + jnp.transpose(edge_bias.astype(compute_dtype), (2, 0, 1)) valid = (idx[None, :] <= idx[:, None]) & mask_bool[None, :] logits = jnp.where( valid[None, :, :], logits, jnp.asarray(-1.0e30, dtype=compute_dtype), ) alpha = jax.nn.softmax(logits, axis=-1) out = jnp.einsum("hij,jhd->ihd", alpha, v_c).astype(dtype) delta = self._dense_no_bias( w_o, self._collapse_heavy_heads(out), tag_id="route.tree_prefix.level.o", kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) proposal_attn = query_mask * delta x = _tree_ngpt_residual( nodes, proposal_attn, self.alpha_tree_level_attn, max_gain=self.tree_ngpt_alpha_max, tag_id="", update_mask=mask, kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) y = self._ln( ffn_ln_s, x, tag_id="route.tree_prefix.level.ffn_ln", kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) y = self._dense( ffn_w1, ffn_b1, y, tag_id="route.tree_prefix.level.ffn1", kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) y = fused_silu(y) y = self._dense( ffn_w2, ffn_b2, y, tag_id="route.tree_prefix.level.ffn2", kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) proposal_ffn = query_mask * y x = _tree_ngpt_residual( x, proposal_ffn, self.alpha_tree_level_ffn, max_gain=self.tree_ngpt_alpha_max, tag_id="", update_mask=mask, kfac_structural_mask=node_structural_mask, kfac_scan_shared=True, kfac_repeat_ndim=1, ) return x, edges def _incremental_edge_parent_vector( self, e00, e01, e10, e11, c0, c1, d0, d1, active, ): skip = jnp.asarray(0.25, dtype=e00.dtype) * (e00 + e01 + e10 + e11) proposal = jax.vmap( lambda x00, x01, x10, x11, xa, xb, ya, yb, keep: self.tree_edge_merge( x00, x01, x10, x11, xa, xb, ya, yb, kfac_structural_mask=keep, kfac_scan_shared=True, ) )(e00, e01, e10, e11, c0, c1, d0, d1, active) merged = self.tree_edge_merge.apply_skip( skip, proposal, kfac_structural_mask=active, kfac_scan_shared=True, ) merged = _tree_sphere(merged) return jnp.where(active[:, None], merged, jnp.zeros_like(merged)) def _apply_tree_level_attention_append( self, raw_nodes, edge_pre, active, row, b_cache, *, edge_row, edge_col, sequence_axis_name=None, sequence_mesh=None, ): edge_row, b_cache = self.tree_level_fwl.append_causal_row( edge_pre, raw_nodes, active, row, b_cache, edge_row=edge_row, edge_col=edge_col, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) edge_row = jnp.where( active[:, None], _tree_sphere(edge_row), edge_row, ) ( ln_s, w_qkv, w_o, edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2, ffn_ln_s, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ) = self.tree_level_layer.as_tuple() n = raw_nodes.shape[0] dtype = raw_nodes.dtype idx = jnp.arange(n, dtype=jnp.int32) x_ln = self._ln( ln_s, raw_nodes, tag_id="route.tree_prefix.level.ln", kfac_structural_mask=active, kfac_scan_shared=True, kfac_repeat_ndim=1, ) qkv = self._dense_no_bias( w_qkv, x_ln, tag_id="route.tree_prefix.level.qkv", kfac_structural_mask=active, kfac_scan_shared=True, kfac_repeat_ndim=1, ).reshape(n, 3, self.n_heads_kernel, self.d_head) q = qkv[row, 0] k = qkv[:, 1] v = qkv[:, 2] edge_bias = self._tree_prefix_edge_bias( edge_row, (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), prefix="level", kfac_structural_mask=active, kfac_scan_shared=True, kfac_repeat_ndim=1, ) pos_bias = lca_alibi_bias( jnp.asarray([row], dtype=jnp.int32), idx, lca_fixed_slopes(self.n_heads_kernel, dtype=dtype), )[:, 0, :] edge_bias = edge_bias + jnp.transpose(pos_bias, (1, 0)) logits = jnp.einsum("hd,jhd->hj", q, k) logits = logits / jnp.sqrt(jnp.asarray(self.d_head, dtype=dtype)) logits = logits + jnp.transpose(edge_bias, (1, 0)) logits = jnp.where( active[None, :], logits, jnp.asarray(-1.0e30, dtype=dtype), ) alpha = jax.nn.softmax(logits, axis=-1) out = jnp.einsum("hj,jhd->hd", alpha, v) delta = self._dense_no_bias( w_o, self._collapse_heavy_heads(out), tag_id="route.tree_prefix.level.o", kfac_structural_mask=jnp.asarray(True), kfac_scan_shared=True, kfac_repeat_ndim=0, ) raw_row = raw_nodes[row] x = _tree_ngpt_residual( raw_row, delta, self.alpha_tree_level_attn, max_gain=self.tree_ngpt_alpha_max, tag_id="", update_mask=jnp.asarray(True), kfac_structural_mask=jnp.asarray(True), kfac_scan_shared=True, kfac_repeat_ndim=0, ) y = self._ln( ffn_ln_s, x, tag_id="route.tree_prefix.level.ffn_ln", kfac_structural_mask=jnp.asarray(True), kfac_scan_shared=True, kfac_repeat_ndim=0, ) y = self._dense( ffn_w1, ffn_b1, y, tag_id="route.tree_prefix.level.ffn1", kfac_structural_mask=jnp.asarray(True), kfac_scan_shared=True, kfac_repeat_ndim=0, ) y = fused_silu(y) y = self._dense( ffn_w2, ffn_b2, y, tag_id="route.tree_prefix.level.ffn2", kfac_structural_mask=jnp.asarray(True), kfac_scan_shared=True, kfac_repeat_ndim=0, ) x = _tree_ngpt_residual( x, y, self.alpha_tree_level_ffn, max_gain=self.tree_ngpt_alpha_max, tag_id="", update_mask=jnp.asarray(True), kfac_structural_mask=jnp.asarray(True), kfac_scan_shared=True, kfac_repeat_ndim=0, ) x = _tree_sphere(x) return x, edge_row, edge_col, b_cache def _incremental_tree_append( self, state, leaf, chosen, t, prefix_ids, edge, mask, *, edge_transpose=None, g=None, sequence_axis_name=None, sequence_mesh=None, ): nodes_raw, nodes_post, level_states = state n = mask.shape[0] n_pad = nodes_raw.shape[1] depth = nodes_raw.shape[0] - 1 dtype = leaf.dtype append = mask[t].astype(bool) leaf_node = _tree_sphere(leaf) nodes_raw = nodes_raw.at[0, t].set( jnp.where(append, leaf_node, jnp.zeros_like(leaf_node)) ) nodes_post = nodes_post.at[0, t].set( jnp.where(append, leaf_node, jnp.zeros_like(leaf_node)) ) if depth == 0: return nodes_raw, nodes_post, level_states mask_pad = jnp.pad(mask.astype(dtype), (0, n_pad - n)) route_pad = jnp.pad( prefix_ids.astype(jnp.int32), (0, n_pad - n), ) clock_depth = self._tree_clock_depth_from_mask(mask) depth_features = _tree_ngpt_level_counts( mask_pad, n_pad // 2, depth, dtype, feature_n_levels=clock_depth, ) clock_state = mask_pad clock_pair_bases = [] fixed_pairs = n_pad // 2 for _level in range(depth): clock_pairs = clock_state.reshape(fixed_pairs, 2) clock_parent = ( clock_pairs[:, 0] + clock_pairs[:, 1] - clock_pairs[:, 0] * clock_pairs[:, 1] ) clock_pair_bases.append( jnp.maximum( jnp.sum(clock_parent.astype(jnp.int32)), jnp.asarray(2, dtype=jnp.int32), ) ) clock_state = jnp.concatenate( [clock_parent, jnp.zeros_like(clock_parent)], axis=0, ) g_projected = self.tree_merge.project_global( g, append & (t > 0), ) if sequence_axis_name is not None: from jax.sharding import NamedSharding, PartitionSpec as P lanes = ( int(sequence_mesh.shape[sequence_axis_name]) if sequence_mesh is not None else 1 ) def constrain_nodes(value): spec = P(None, sequence_axis_name, None) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) def constrain_level(level_state): width_local = level_state[0].shape[0] row_axis = ( sequence_axis_name if width_local >= lanes and width_local % lanes == 0 else None ) spec = P(row_axis, None, None) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return tuple( jax.lax.with_sharding_constraint(value, spec) for value in level_state ) def constrain_row_value(value): spec = P(None, None) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) def constrain_column_value(value): row_axis = ( sequence_axis_name if value.shape[0] >= lanes and value.shape[0] % lanes == 0 else None ) spec = P(row_axis, None) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) else: constrain_nodes = lambda value: value constrain_level = lambda level_state: level_state constrain_row_value = lambda value: value constrain_column_value = lambda value: value levels_mut = list(level_states) for level in range(depth): width = n_pad >> (level + 1) block = 1 << (level + 1) create = append & (((t + 1) % block) == 0) parent = (t + 1) // block - 1 lower_post_edges = None if level == 0 else levels_mut[level - 1][1] def do_create(operand): raw_all, post_all, level_state = operand edge_pre, edge_post, b_cache = level_state q = jnp.arange(width, dtype=jnp.int32) q0 = 2 * q q1 = q0 + 1 p0 = 2 * parent p1 = p0 + 1 children = post_all[level] left = children[p0] right = children[p1] if level == 0: source0 = route_pad[p0] source1 = route_pad[p1] row0 = self._tree_pair_message_row( edge, source0, mask, edge_transpose=edge_transpose, )[route_pad] row1 = self._tree_pair_message_row( edge, source1, mask, edge_transpose=edge_transpose, )[route_pad] col0 = self._tree_pair_message_column( edge, source0, mask, edge_transpose=edge_transpose, )[route_pad] col1 = self._tree_pair_message_column( edge, source1, mask, edge_transpose=edge_transpose, )[route_pad] row0 = _tree_sphere(row0) row1 = _tree_sphere(row1) col0 = _tree_sphere(col0) col1 = _tree_sphere(col1) row_cells = ( row0[q0], row0[q1], row1[q0], row1[q1], ) col_cells = ( col0[q0], col1[q0], col0[q1], col1[q1], ) sibling_lr = row0[p1] sibling_rl = row1[p0] else: assert lower_post_edges is not None lower_row0 = _square_row_by_reduction( lower_post_edges, p0, ) lower_row1 = _square_row_by_reduction( lower_post_edges, p1, ) lower_col0 = _square_column_local( lower_post_edges, p0, ) lower_col1 = _square_column_local( lower_post_edges, p1, ) row_cells = ( lower_row0[q0], lower_row0[q1], lower_row1[q0], lower_row1[q1], ) col_cells = ( lower_col0[q0], lower_col1[q0], lower_col0[q1], lower_col1[q1], ) sibling_lr = lower_row0[p1] sibling_rl = lower_row1[p0] depth_row = ( None if depth_features is None else depth_features[level, parent][None, :] ) merged, _valid, _genuine = self.tree_merge( left[None, :], right[None, :], jnp.ones((1,), dtype=dtype), jnp.ones((1,), dtype=dtype), sibling_lr[None, :], sibling_rl[None, :], jnp.asarray(level, dtype=jnp.int32), jnp.asarray([parent], dtype=jnp.int32), clock_pair_bases[level], clock_depth, depth_feats=depth_row, g=g, g_structural_mask=jnp.asarray(True), g_projected=g_projected, ) raw_parent = merged[0] raw_all = raw_all.at[level + 1, parent].set(raw_parent) active = q <= parent new_left = jnp.broadcast_to(left, (width, self.d_model)) new_right = jnp.broadcast_to(right, (width, self.d_model)) other_left = children[q0] other_right = children[q1] parent_row = self._incremental_edge_parent_vector( *row_cells, new_left, new_right, other_left, other_right, active, ) parent_col = self._incremental_edge_parent_vector( *col_cells, other_left, other_right, new_left, new_right, active, ) parent_row = constrain_row_value(parent_row) parent_col = constrain_column_value(parent_col) edge_pre = _replace_square_row_column( edge_pre, parent, parent_row, parent_col, ) post_parent, post_row, post_col, b_cache = ( self._apply_tree_level_attention_append( raw_all[level + 1, :width], edge_pre, active, parent, b_cache, edge_row=parent_row, edge_col=parent_col, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) ) post_all = post_all.at[level + 1, parent].set(post_parent) post_row = constrain_row_value(post_row) post_col = constrain_column_value(post_col) edge_post = _replace_square_row_column( edge_post, parent, post_row, post_col, ) return raw_all, post_all, (edge_pre, edge_post, b_cache) nodes_raw, nodes_post, levels_mut[level] = jax.lax.cond( create, do_create, lambda operand: operand, (nodes_raw, nodes_post, levels_mut[level]), ) nodes_raw = constrain_nodes(nodes_raw) nodes_post = constrain_nodes(nodes_post) levels_mut[level] = constrain_level(levels_mut[level]) return nodes_raw, nodes_post, tuple(levels_mut) def _tree_prefix_scan(self, seq, mask, pair_route, *, clock_mask=None, g=None): n = seq.shape[0] n_pad = _route_next_pow2(n) depth = n_pad.bit_length() - 1 dtype = seq.dtype pad = n_pad - n clock_mask = mask if clock_mask is None else clock_mask clock_depth = self._tree_clock_depth_from_mask(clock_mask) nodes0 = jnp.pad(seq, ((0, pad), (0, 0))) valid0 = jnp.pad(mask.astype(dtype), (0, pad)) nodes0 = jnp.where( valid0.astype(bool)[:, None], _tree_sphere(nodes0), jnp.zeros_like(nodes0), ) clock_valid0 = jnp.pad(clock_mask.astype(dtype), (0, pad)) tree_has_merge = jnp.sum(valid0.astype(jnp.int32)) > 1 edge0 = jnp.pad(pair_route, ((0, pad), (0, pad), (0, 0))) edge0_active = valid0.astype(bool)[:, None] & valid0.astype(bool)[None, :] edge0 = jnp.where( edge0_active[..., None], _tree_sphere(edge0), jnp.zeros_like(edge0), ) if depth == 0: levels = nodes0[None, :, :] valids = valid0[None, :] edges = edge0[None, :, :, :] scan_nodes = jnp.zeros((0,) + nodes0.shape, dtype=dtype) return levels, valids, valids, edges, scan_nodes n_pairs = n_pad // 2 g_projected = self.tree_merge.project_global(g, tree_has_merge) def _split(x): xr = x.reshape((n_pairs, 2) + x.shape[1:]) return xr[:, 0], xr[:, 1] def _zpad(x): return jnp.concatenate([x, jnp.zeros_like(x)], axis=0) def _zpad_edge(x): pad_n = n_pad - n_pairs return jnp.pad(x, ((0, pad_n), (0, pad_n), (0, 0))) pidx = jnp.arange(n_pairs, dtype=jnp.int32) def body(state, xs_lv): level_idx, depth_feats_lv = xs_lv nodes, valid, edge_state, clock_valid = state left, right = _split(nodes) left_m, right_m = _split(valid) clock_left_m, clock_right_m = _split(clock_valid) clock_pair_active = ( clock_left_m + clock_right_m - clock_left_m * clock_right_m ) pair_base = jnp.maximum( jnp.sum(clock_pair_active.astype(jnp.int32)), jnp.asarray(2, dtype=jnp.int32), ) e_rs = edge_state.reshape(n_pairs, 2, n_pairs, 2, self.d_model) merged, out_mask, genuine = self.tree_merge( left, right, left_m, right_m, e_rs[pidx, 0, pidx, 1, :], e_rs[pidx, 1, pidx, 0, :], level_idx, pidx, pair_base, clock_depth, depth_feats=depth_feats_lv, g=g, g_structural_mask=tree_has_merge, g_projected=g_projected, ) valid_pair = valid.reshape(n_pairs, 2).astype(dtype) weights = valid_pair[:, :, None, None] * valid_pair[None, None, :, :] cell_count = jnp.sum(weights, axis=(1, 3)) denom = jnp.maximum( cell_count, jnp.asarray(1.0, dtype=dtype), ) edge_parent = ( jnp.sum(e_rs * weights[..., None], axis=(1, 3)) / denom[..., None] ) e00 = e_rs[:, 0, :, 0, :] e01 = e_rs[:, 0, :, 1, :] e10 = e_rs[:, 1, :, 0, :] e11 = e_rs[:, 1, :, 1, :] edge_keep = (genuine[:, None] * genuine[None, :]).astype(bool) def _edge_row(e0, e1, e2, e3, c0, c1, keep_row): return jax.vmap( lambda x0, x1, x2, x3, d0, d1, keep: self.tree_edge_merge( x0, x1, x2, x3, c0, c1, d0, d1, kfac_structural_mask=keep, kfac_scan_shared=True, ) )(e0, e1, e2, e3, left, right, keep_row) edge_proposal = jax.vmap(_edge_row)( e00, e01, e10, e11, left, right, edge_keep, ) edge_updated = self.tree_edge_merge.apply_skip( edge_parent, edge_proposal, kfac_structural_mask=edge_keep, kfac_scan_shared=True, ) edge_parent = jnp.where( edge_keep[..., None], edge_updated, edge_parent, ) edge_parent = jnp.where( (cell_count > 1)[..., None], _tree_sphere(edge_parent), edge_parent, ) merged_skip = merged edge_skip = edge_parent merged, edge_parent = self._apply_tree_level_attention( merged, edge_parent, genuine, level_idx, ) merged = jnp.where( genuine.astype(bool)[:, None], _tree_sphere(merged), merged_skip, ) edge_update_mask = ( genuine.astype(bool)[:, None] & genuine.astype(bool)[None, :] ) idx = jnp.arange(n_pairs, dtype=jnp.int32) edge_update_mask = edge_update_mask & (idx[None, :] <= idx[:, None]) edge_parent = jnp.where( edge_update_mask[..., None], _tree_sphere(edge_parent), edge_skip, ) next_state = ( _zpad(merged), _zpad(out_mask), _zpad_edge(edge_parent), _zpad(clock_pair_active), ) ys = (next_state[0], next_state[1], _zpad(genuine), next_state[2]) return next_state, ys depth_feat_levels = _tree_ngpt_level_counts( clock_valid0, n_pairs, depth, dtype, ) (_nodes, _valid, _edge, _clock_valid), ys = jax.lax.scan( body, (nodes0, valid0, edge0, clock_valid0), (jnp.arange(depth, dtype=jnp.int32), depth_feat_levels), ) nodes_y, valid_y, genuine_y, edge_y = ys tree_levels = jnp.concatenate([nodes0[None, :, :], nodes_y], axis=0) valid_levels = jnp.concatenate([valid0[None, :], valid_y], axis=0) genuine_levels = jnp.concatenate([valid0[None, :], genuine_y], axis=0) edge_levels = jnp.concatenate([edge0[None, :, :, :], edge_y], axis=0) return tree_levels, valid_levels, genuine_levels, edge_levels, nodes_y def _source_edge_levels(self, pair_msg, mask): n = pair_msg.shape[0] n_dst = pair_msg.shape[1] n_pad = _route_next_pow2(n) depth = n_pad.bit_length() - 1 dtype = pair_msg.dtype pad = n_pad - n edge_state = jnp.pad(pair_msg, ((0, pad), (0, 0), (0, 0))) valid = jnp.pad(mask.astype(dtype), (0, pad)) levels = [edge_state] n_pairs = n_pad // 2 for _level in range(depth): e_rs = edge_state.reshape(n_pairs, 2, n_dst, self.d_model) v_rs = valid.reshape(n_pairs, 2) weights = v_rs[:, :, None, None] denom = jnp.maximum( jnp.sum(v_rs, axis=1), jnp.asarray(1.0, dtype=dtype), ) parent = jnp.sum(e_rs * weights, axis=1) / denom[:, None, None] valid_parent = v_rs[:, 0] + v_rs[:, 1] - v_rs[:, 0] * v_rs[:, 1] edge_state = jnp.concatenate([parent, jnp.zeros_like(parent)], axis=0) valid = jnp.concatenate( [valid_parent, jnp.zeros_like(valid_parent)], axis=0 ) levels.append(edge_state) return jnp.stack(levels, axis=0) def _prefix_cover(self, n: int): n_pad = _route_next_pow2(n) depth = n_pad.bit_length() - 1 if depth == 0: return ( jnp.zeros((n, 0), dtype=jnp.int32), jnp.zeros((n, 0), dtype=jnp.int32), jnp.zeros((n, 0), dtype=bool), ) t = jnp.arange(n, dtype=jnp.int32) start = jnp.zeros((n,), dtype=jnp.int32) levels = [] nodes = [] valids = [] for bit in range(depth - 1, -1, -1): take = (jnp.right_shift(t, bit) & 1) == 1 node = jnp.right_shift(start, bit) levels.append(jnp.where(take, jnp.asarray(bit, jnp.int32), 0)) nodes.append(jnp.where(take, node, 0)) valids.append(take) start = start + jnp.where(take, jnp.asarray(1 << bit, jnp.int32), 0) level_arr = jnp.stack(levels, axis=1) node_arr = jnp.stack(nodes, axis=1) valid_arr = jnp.stack(valids, axis=1) prefix_width = _route_next_pow2(depth) pad = prefix_width - depth if pad: level_arr = jnp.pad(level_arr, ((0, 0), (0, pad))) node_arr = jnp.pad(node_arr, ((0, 0), (0, pad))) valid_arr = jnp.pad(valid_arr, ((0, 0), (0, pad)), constant_values=False) return level_arr, node_arr, valid_arr def _segment_weights(self, cover_level, cover_node, cover_valid, mask, dtype): n = mask.shape[0] idx = jnp.arange(n, dtype=jnp.int32) leaf_node = jnp.right_shift(idx[None, None, :], cover_level[..., None]) member = ( (leaf_node == cover_node[..., None]) & cover_valid[..., None] & mask.astype(bool)[None, None, :] ) weights = member.astype(dtype) denom = jnp.maximum( jnp.sum(weights, axis=-1, keepdims=True), jnp.asarray(1.0, dtype=dtype), ) return weights / denom def _tree_prefix_context_all( self, seq, mask, pair_route, pair_to_nodes, route_ids, g=None, ): n = seq.shape[0] dtype = seq.dtype cover_level, cover_node, cover_valid = self._prefix_cover(n) depth = cover_level.shape[1] if depth == 0: return ( jnp.zeros((n, 0, self.d_model), dtype=dtype), jnp.zeros((n, n, 0, self.d_model), dtype=dtype), jnp.zeros((n, 0, 0, self.d_model), dtype=dtype), jnp.zeros((n, 0), dtype=bool), None, ) tree_levels, _valid_levels, genuine_levels, _edge_levels, _ys = ( self._tree_prefix_scan(seq, mask, pair_route, g=g) ) source_to_nodes = self._source_edge_levels(pair_to_nodes, mask) prefix_nodes = tree_levels[cover_level, cover_node] prefix_mask = ( cover_valid & (genuine_levels[cover_level, cover_node] > 0) & mask.astype(bool)[:, None] ) source_nodes = source_to_nodes[cover_level, cover_node] cand_prefix_edge = jnp.transpose(source_nodes, (0, 2, 1, 3)) source_route = jnp.take(source_nodes, route_ids, axis=-2) dst_weights = self._segment_weights( cover_level, cover_node, cover_valid, mask, dtype, ) prefix_prefix_edge = jnp.einsum("tlsd,tms->tlmd", source_route, dst_weights) g_rows = self._causal_prefix_g(g, prefix_nodes, prefix_mask) return (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_rows) def _causal_prefix_g(self, g, prefix_nodes, prefix_mask): if prefix_nodes.ndim == 2: update_active = jnp.any(prefix_mask.astype(bool)) return self.g_step_update( g, self.g_step_pool( g, prefix_nodes, prefix_mask.astype(prefix_nodes.dtype), kfac_structural_mask=prefix_mask.astype(bool), kfac_update_mask=update_active, kfac_repeat_ndim=1, ), update_mask=update_active, kfac_structural_mask=update_active, kfac_repeat_ndim=0, ) pool_query_active = jnp.any(prefix_mask.astype(bool)) return jax.vmap( lambda pn, pm: self.g_step_update( g, self.g_step_pool( g, pn, pm, kfac_structural_mask=pm.astype(bool), kfac_update_mask=pool_query_active, kfac_repeat_ndim=2, ), update_mask=jnp.any(pm.astype(bool)), kfac_structural_mask=jnp.any(pm.astype(bool)), kfac_g_structural_mask=pool_query_active, kfac_repeat_ndim=1, ) )(prefix_nodes, prefix_mask.astype(prefix_nodes.dtype)) def _tree_prefix_context_row( self, seq, mask, pair_route, pair_to_nodes, route_ids, t, *, clock_mask=None, g=None, source_edge_frontier=None, source_edge_counts=None, sequence_axis_name=None, sequence_mesh=None, ): n = seq.shape[0] dtype = seq.dtype cover_level, cover_node, cover_valid = self._prefix_cover(n) depth = cover_level.shape[1] if depth == 0: return ( jnp.zeros((0, self.d_model), dtype=dtype), jnp.zeros((n, 0, self.d_model), dtype=dtype), jnp.zeros((0, 0, self.d_model), dtype=dtype), jnp.zeros((0,), dtype=bool), None, ) tree_levels, _valid_levels, genuine_levels, _edge_levels, _ys = ( self._tree_prefix_scan(seq, mask, pair_route, clock_mask=clock_mask, g=g) ) cl = cover_level[t] cn = cover_node[t] cv = cover_valid[t] prefix_nodes = tree_levels[cl, cn] row_active = ( mask[t].astype(bool) if clock_mask is None else clock_mask[t].astype(bool) ) prefix_mask = cv & (genuine_levels[cl, cn] > 0) & row_active if source_edge_frontier is None: source_to_nodes = self._source_edge_levels(pair_to_nodes, mask) source_nodes = source_to_nodes[cl, cn] else: if source_edge_counts is None: raise ValueError( "source_edge_counts is required with source_edge_frontier" ) source_sums = source_edge_frontier[cl] source_counts = source_edge_counts[cl] source_nodes = ( source_sums / jnp.maximum( source_counts, jnp.asarray(1.0, dtype=dtype), )[:, None, None] ) source_nodes = jnp.where( (cv & (source_counts > 0))[:, None, None], source_nodes, jnp.zeros_like(source_nodes), ) dst_weights = self._segment_weights( cl[None, :], cn[None, :], cv[None, :], mask, dtype, )[0] weights_by_candidate = ( jnp.zeros_like(dst_weights).at[:, route_ids].add(dst_weights) ) if sequence_axis_name is not None: from jax.sharding import NamedSharding, PartitionSpec as P def _sharding(*axes): spec = P(*axes) return ( NamedSharding(sequence_mesh, spec) if sequence_mesh is not None else spec ) source_nodes = jax.lax.with_sharding_constraint( source_nodes, _sharding(None, sequence_axis_name, None), ) weights_by_candidate = jax.lax.with_sharding_constraint( weights_by_candidate, _sharding(None, None), ) cand_prefix_edge = jnp.transpose(source_nodes, (1, 0, 2)) prefix_prefix_edge = jnp.einsum( "lcd,mc->lmd", source_nodes, weights_by_candidate, ) if sequence_axis_name is not None: cand_prefix_edge = jax.lax.with_sharding_constraint( cand_prefix_edge, _sharding(sequence_axis_name, None, None), ) prefix_prefix_edge = jax.lax.with_sharding_constraint( prefix_prefix_edge, _sharding(None, None, None), ) g_row = self._causal_prefix_g(g, prefix_nodes, prefix_mask) return (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row) def _incremental_tree_state(self, n: int, dtype): n_pad = _route_next_pow2(n) depth = n_pad.bit_length() - 1 nodes_raw = jnp.zeros((depth + 1, n_pad, self.d_model), dtype=dtype) nodes_post = jnp.zeros_like(nodes_raw) channels = self.tree_level_fwl.two_hop_channels levels = [] for level in range(depth): width = n_pad >> (level + 1) edge_pre = jnp.zeros( (width, width, self.d_model), dtype=dtype, ) edge_post = jnp.zeros_like(edge_pre) b_cache = jnp.zeros((width, width, channels), dtype=dtype) levels.append((edge_pre, edge_post, b_cache)) return nodes_raw, nodes_post, tuple(levels) def _tree_prefix_context_row_incremental( self, nodes_post, mask, route_ids, t, *, g=None, source_edge_frontier, source_edge_counts, sequence_axis_name=None, sequence_mesh=None, ): n = mask.shape[0] dtype = nodes_post.dtype cover_level, cover_node, cover_valid = self._prefix_cover(n) cl = cover_level[t] cn = cover_node[t] cv = cover_valid[t] prefix_nodes = nodes_post[cl, cn] prefix_mask = cv & mask[t].astype(bool) source_sums = source_edge_frontier[cl] source_counts = source_edge_counts[cl] source_nodes = ( source_sums / jnp.maximum( source_counts, jnp.asarray(1.0, dtype=dtype), )[:, None, None] ) source_nodes = jnp.where( (cv & (source_counts > 0))[:, None, None], source_nodes, jnp.zeros_like(source_nodes), ) dst_weights = self._segment_weights( cl[None, :], cn[None, :], cv[None, :], mask, dtype, )[0] weights_by_candidate = ( jnp.zeros_like(dst_weights).at[:, route_ids].add(dst_weights) ) if sequence_axis_name is not None: from jax.sharding import NamedSharding, PartitionSpec as P def _sharding(*axes): spec = P(*axes) return ( NamedSharding(sequence_mesh, spec) if sequence_mesh is not None else spec ) source_nodes = jax.lax.with_sharding_constraint( source_nodes, _sharding(None, sequence_axis_name, None), ) weights_by_candidate = jax.lax.with_sharding_constraint( weights_by_candidate, _sharding(None, None), ) cand_prefix_edge = jnp.transpose(source_nodes, (1, 0, 2)) prefix_prefix_edge = jnp.einsum( "lcd,mc->lmd", source_nodes, weights_by_candidate, ) if sequence_axis_name is not None: cand_prefix_edge = jax.lax.with_sharding_constraint( cand_prefix_edge, _sharding(sequence_axis_name, None, None), ) prefix_prefix_edge = jax.lax.with_sharding_constraint( prefix_prefix_edge, _sharding(None, None, None), ) g_row = self._causal_prefix_g(g, prefix_nodes, prefix_mask) return ( prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row, ) def _tree_prefix_edge_bias( self, edge_msg, params, *, prefix: str, kfac_structural_mask, kfac_scan_shared: bool = False, kfac_repeat_ndim: int = 2, kfac_context_primal_reused_over_walkers: bool = False, ): ln_s, w1, b1, w2, b2 = params x = self._ln( ln_s, edge_msg, tag_id=f"route.tree_prefix.{prefix}.edge_ln", kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) x = self._dense( w1, b1, x, tag_id=f"route.tree_prefix.{prefix}.edge_bias1", kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) x = fused_silu(x) return self._dense( w2, b2, x, tag_id=f"route.tree_prefix.{prefix}.edge_bias2", kfac_structural_mask=kfac_structural_mask, kfac_scan_shared=kfac_scan_shared, kfac_repeat_ndim=kfac_repeat_ndim, kfac_context_primal_reused_over_walkers=( kfac_context_primal_reused_over_walkers ), ) def _tree_query_seed(self, query_global, route_pos, mask, dtype): pos = jnp.asarray(route_pos, dtype=jnp.int32) target_shape = pos.shape + (self.d_model,) if query_global is None: seed = jnp.zeros(target_shape, dtype=dtype) else: seed = jnp.broadcast_to( jnp.asarray(query_global, dtype=dtype), target_shape, ) seed = seed + self._route_position_embedding( pos, dtype, mask=mask, ) return seed def _tree_prefix_layer( self, x, prefix_edges, token_mask, attention_mask, token_structural_mask, pair_structural_mask, params, *, impl, g_projection, ): ( ln_s, w_qkv, w_o, edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2, ffn_ln_s, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ) = params context_reuse = True token_kfac = dict( kfac_structural_mask=token_structural_mask, kfac_repeat_ndim=2, kfac_context_primal_reused_over_walkers=context_reuse, ) bsz, n_tokens = x.shape[:2] residual_mask = token_mask.astype(x.dtype)[..., None] x_ln = self._ln( ln_s, x, tag_id="route.tree_prefix.graph.ln", **token_kfac, ) qkv = self._dense_no_bias( w_qkv, x_ln, tag_id="route.tree_prefix.graph.qkv", **token_kfac, ).reshape(bsz, n_tokens, 3, self.n_heads_kernel, self.d_head) q = qkv[:, :, 0] k = qkv[:, :, 1] v = qkv[:, :, 2] edge_bias = self._tree_prefix_edge_bias( prefix_edges, (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), prefix="graph", kfac_structural_mask=pair_structural_mask, kfac_repeat_ndim=3, kfac_context_primal_reused_over_walkers=context_reuse, ) out = self._route_attention( q, k, v, edge_bias, token_mask, impl=impl, attention_mask=attention_mask, ) delta = self._dense_no_bias( w_o, self._collapse_heavy_heads(out), tag_id="route.tree_prefix.graph.o", **token_kfac, ) x = x + residual_mask * self.tree_prefix_residual_gain * delta y = self._ln( ffn_ln_s, x, tag_id="route.tree_prefix.graph.ffn_ln", **token_kfac, ) if g_projection is not None: y = y + g_projection y = self._dense( ffn_w1, ffn_b1, y, tag_id="route.tree_prefix.graph.ffn1", **token_kfac, ) y = fused_silu(y) y = self._dense( ffn_w2, ffn_b2, y, tag_id="route.tree_prefix.graph.ffn2", **token_kfac, ) return x + residual_mask * self.tree_prefix_residual_gain * y def _tree_candidate_layer( self, cand, prefix_nodes, cand_prefix_edge, prefix_mask, cand_mask, candidate_structural_mask, prefix_structural_mask, cross_pair_structural_mask, params, *, impl, g_projection, sequence_axis_name=None, sequence_mesh=None, ): ( cand_ln_s, pref_ln_s, cand_w_qv, pref_w_kv, w_o, edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2, ffn_ln_s, ffn_w1, ffn_b1, ffn_w2, ffn_b2, ) = params context_reuse = True candidate_kfac = dict( kfac_structural_mask=candidate_structural_mask, kfac_repeat_ndim=2, kfac_context_primal_reused_over_walkers=context_reuse, ) prefix_kfac = dict( kfac_structural_mask=prefix_structural_mask, kfac_repeat_ndim=2, kfac_context_primal_reused_over_walkers=context_reuse, ) bsz, n_cand = cand.shape[:2] n_pref = prefix_nodes.shape[1] def _seq_constraint(value, *axes): if sequence_axis_name is None: return value from jax.sharding import NamedSharding, PartitionSpec as P spec = P(*axes) if sequence_mesh is not None: spec = NamedSharding(sequence_mesh, spec) return jax.lax.with_sharding_constraint(value, spec) cand = _seq_constraint(cand, None, sequence_axis_name, None) cand_ln = self._ln( cand_ln_s, cand, tag_id="route.tree_prefix.candidate.ln", **candidate_kfac, ) if sequence_axis_name is None: qv = self._dense_no_bias( cand_w_qv, cand_ln, tag_id="route.tree_prefix.candidate.qv", **candidate_kfac, ).reshape(bsz, n_cand, 2, self.n_heads_kernel, self.d_head) q = qv[:, :, 0] v_self = qv[:, :, 1] else: d_qv = self.n_heads_kernel * self.d_head q = jnp.matmul(cand_ln, cand_w_qv[:, :d_qv]).reshape( bsz, n_cand, self.n_heads_kernel, self.d_head, ) v_self = jnp.matmul(cand_ln, cand_w_qv[:, d_qv:]).reshape( bsz, n_cand, self.n_heads_kernel, self.d_head, ) q = _seq_constraint(q, None, sequence_axis_name, None, None) v_self = _seq_constraint( v_self, None, sequence_axis_name, None, None, ) pref_ln = self._ln( pref_ln_s, prefix_nodes, tag_id="route.tree_prefix.candidate.prefix_ln", **prefix_kfac, ) kv = self._dense_no_bias( pref_w_kv, pref_ln, tag_id="route.tree_prefix.candidate.kv", **prefix_kfac, ).reshape(bsz, n_pref, 2, self.n_heads_kernel, self.d_head) kv = _seq_constraint(kv, None, None, None, None, None) k = kv[:, :, 0] v = kv[:, :, 1] edge_bias = self._tree_prefix_edge_bias( cand_prefix_edge, (edge_ln_s, edge_w1, edge_b1, edge_w2, edge_b2), prefix="candidate", kfac_structural_mask=cross_pair_structural_mask, kfac_repeat_ndim=3, kfac_context_primal_reused_over_walkers=context_reuse, ) edge_bias = _seq_constraint( edge_bias, None, sequence_axis_name, None, None, ) out = self._route_attention( q, k, v, edge_bias, prefix_mask, impl=impl, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) out = _seq_constraint(out, None, sequence_axis_name, None, None) delta = self._dense_no_bias( w_o, self._collapse_heavy_heads(out), tag_id="route.tree_prefix.candidate.o", **candidate_kfac, ) cand = cand + cand_mask * self.tree_candidate_residual_gain * delta y = self._ln( ffn_ln_s, cand, tag_id="route.tree_prefix.candidate.ffn_ln", **candidate_kfac, ) if g_projection is not None: y = y + g_projection y = self._dense( ffn_w1, ffn_b1, y, tag_id="route.tree_prefix.candidate.ffn1", **candidate_kfac, ) y = fused_silu(y) y = self._dense( ffn_w2, ffn_b2, y, tag_id="route.tree_prefix.candidate.ffn2", **candidate_kfac, ) return cand + cand_mask * self.tree_candidate_residual_gain * y def _apply_tree_prefix_layers( self, prefix_nodes, prefix_edges, prefix_mask, g=None, query_seed=None, row_mask=None, ): wants_query = query_seed is not None single = prefix_nodes.ndim == 2 if single: prefix_nodes = prefix_nodes[None, :, :] prefix_edges = prefix_edges[None, :, :, :] prefix_mask = prefix_mask[None, :] if wants_query: query_seed = query_seed[None, :] if row_mask is not None: row_mask = jnp.asarray(row_mask, dtype=bool).reshape(1) if wants_query: assert query_seed is not None query_seed = jnp.broadcast_to( query_seed, prefix_nodes.shape[:-2] + (self.d_model,), ) n_pref = prefix_nodes.shape[1] x = jnp.concatenate([prefix_nodes, query_seed[:, None, :]], axis=1) prefix_edges = jnp.pad( prefix_edges, ((0, 0), (0, 1), (0, 1), (0, 0)), ) token_mask = jnp.concatenate( [ prefix_mask.astype(bool), jnp.ones((prefix_mask.shape[0], 1), dtype=bool), ], axis=1, ) token_idx = jnp.arange(n_pref + 1, dtype=jnp.int32) query_row = token_idx == n_pref cover_key = token_idx < n_pref attention_mask = token_mask[:, None, :] & ( query_row[None, :, None] | cover_key[None, None, :] ) else: n_pref = prefix_nodes.shape[1] x = prefix_nodes token_mask = prefix_mask.astype(bool) attention_mask = None row_structural_mask = ( jnp.ones((x.shape[0],), dtype=bool) if row_mask is None else jnp.broadcast_to(jnp.asarray(row_mask, dtype=bool), (x.shape[0],)) ) token_structural_mask = token_mask.astype(bool) & row_structural_mask[:, None] if attention_mask is None: pair_structural_mask = ( token_structural_mask[:, :, None] & token_structural_mask[:, None, :] ) else: pair_structural_mask = ( token_structural_mask[:, :, None] & token_structural_mask[:, None, :] & attention_mask.astype(bool) ) if self.route_tree_prefix_layers == 0: cover = x[:, :n_pref] if not wants_query: return cover[0] if single else cover query = x[:, n_pref] return ( cover[0] if single else cover, query[0] if single else query, ) impl = self._resolve_tree_attn_impl() from hamiltonzero.model.tree import _tagged_dense_no_bias as _tdnb _gg_pref = _tdnb( self.g_prefix_ffn_w, g, tag_id="gladder.route.prefix_fproj", pathway="even", kfac_structural_mask=jnp.any(row_structural_mask), kfac_repeat_ndim=0, kfac_context_primal_reused_over_walkers=True, ).astype(prefix_nodes.dtype) params = self._tree_prefix_layer_params() def apply_one(state, layer): return self._tree_prefix_layer( state, prefix_edges, token_mask, attention_mask, token_structural_mask, pair_structural_mask, layer, impl=impl, g_projection=_gg_pref, ) for layer in params: x = apply_one(x, layer) cover = x[:, :n_pref] if not wants_query: return cover[0] if single else cover query = x[:, n_pref] return ( cover[0] if single else cover, query[0] if single else query, ) def _apply_tree_candidate_layers( self, base, prefix_nodes, cand_prefix_edge, prefix_mask, mask, g_rows=None, candidate_mask=None, sequence_axis_name=None, sequence_mesh=None, ): if self.route_tree_prefix_candidate_layers == 0 or prefix_nodes.shape[-2] == 0: return base single = base.ndim == 2 if single: base = base[None, :, :] prefix_nodes = prefix_nodes[None, :, :] cand_prefix_edge = cand_prefix_edge[None, :, :, :] prefix_mask = prefix_mask[None, :] if candidate_mask is not None: candidate_mask = jnp.asarray(candidate_mask, dtype=bool)[None, :] cand = base impl = self._resolve_tree_attn_impl() cand_mask = mask.astype(cand.dtype)[None, :, None] candidate_structural_mask = ( jnp.broadcast_to(mask.astype(bool), cand.shape[:2]) if candidate_mask is None else jnp.broadcast_to( jnp.asarray(candidate_mask, dtype=bool), cand.shape[:2] ) ) row_structural_mask = jnp.any(candidate_structural_mask, axis=-1) prefix_structural_mask = prefix_mask.astype(bool) & row_structural_mask[:, None] cross_pair_structural_mask = ( candidate_structural_mask[:, :, None] & prefix_structural_mask[:, None, :] ) from hamiltonzero.model.tree import _tagged_dense_no_bias as _tdnb _gg = _tdnb( self.g_cand_ffn_w, g_rows, tag_id="gladder.route.cand_fproj", pathway="even", kfac_structural_mask=( row_structural_mask if jnp.ndim(g_rows) > 1 else jnp.any(row_structural_mask) ), kfac_repeat_ndim=(1 if jnp.ndim(g_rows) > 1 else 0), kfac_context_primal_reused_over_walkers=True, ).astype(cand.dtype) _gg_cand = _gg[..., None, :] if _gg.ndim == cand.ndim - 1 else _gg params = self._tree_candidate_layer_params() def apply_one(state, layer): return self._tree_candidate_layer( state, prefix_nodes, cand_prefix_edge, prefix_mask, cand_mask, candidate_structural_mask, prefix_structural_mask, cross_pair_structural_mask, layer, impl=impl, g_projection=_gg_cand, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) for layer in params: cand = apply_one(cand, layer) return cand[0] if single else cand def _tree_enrich_teacher( self, base, seq, edge, perm, mask, g=None, query_global=None, ): pair_msg = self._tree_pair_messages(edge, mask) pair_route = pair_msg[perm[:, None], perm[None, :]] pair_to_nodes = pair_msg[perm, :] (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_rows) = ( self._tree_prefix_context_all( seq, mask, pair_route, pair_to_nodes, perm, g=g ) ) query_seed = self._tree_query_seed( query_global, jnp.arange(base.shape[0], dtype=jnp.int32), mask, base.dtype, ) idx = jnp.arange(base.shape[0], dtype=jnp.int32) mask_bool = mask.astype(bool) pos_of_node = jnp.zeros((base.shape[0],), dtype=jnp.int32).at[perm].set(idx) candidate_structural_mask = ( mask_bool[:, None] & mask_bool[None, :] & (pos_of_node[None, :] >= idx[:, None]) ) prefix_nodes, query = self._apply_tree_prefix_layers( prefix_nodes, prefix_prefix_edge, prefix_mask, g=g, query_seed=query_seed, row_mask=mask, ) candidate = self._apply_tree_candidate_layers( base, prefix_nodes, cand_prefix_edge, prefix_mask, mask, g_rows=g_rows, candidate_mask=candidate_structural_mask, ) return candidate, query def _apply_tree_prefix_step( self, base, base_cache, prefix_ids, mask, t, pair_msg, g=None, query_global=None, picked=None, source_edge_frontier=None, source_edge_counts=None, raw_edge_for_pair_messages=None, raw_edge_transpose=None, sequence_axis_name=None, sequence_mesh=None, row_permute_fn=None, incremental_tree_state=None, ): idx = jnp.arange(base.shape[0], dtype=jnp.int32) prefix_mask_positions = mask.astype(bool) & (idx < t) if incremental_tree_state is None: pair_route = ( pair_msg[prefix_ids[:, None], prefix_ids[None, :]] if raw_edge_for_pair_messages is None else self._tree_pair_messages_for_route( raw_edge_for_pair_messages, prefix_ids, mask, edge_transpose=raw_edge_transpose, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, row_permute_fn=row_permute_fn, ) ) pair_to_nodes = ( pair_msg[prefix_ids, :] if source_edge_frontier is None else None ) (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row) = ( self._tree_prefix_context_row( base_cache, prefix_mask_positions, pair_route, pair_to_nodes, prefix_ids, t, clock_mask=mask, g=g, source_edge_frontier=source_edge_frontier, source_edge_counts=source_edge_counts, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) ) else: if source_edge_frontier is None or source_edge_counts is None: raise ValueError( "incremental tree context requires source-edge frontiers" ) (prefix_nodes, cand_prefix_edge, prefix_prefix_edge, prefix_mask, g_row) = ( self._tree_prefix_context_row_incremental( incremental_tree_state[1], mask, prefix_ids, t, g=g, source_edge_frontier=source_edge_frontier, source_edge_counts=source_edge_counts, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) ) query_seed = self._tree_query_seed( query_global, t, mask, base.dtype, ) prefix_nodes, query = self._apply_tree_prefix_layers( prefix_nodes, prefix_prefix_edge, prefix_mask, g=g, query_seed=query_seed, row_mask=mask[t], ) candidate = self._apply_tree_candidate_layers( base, prefix_nodes, cand_prefix_edge, prefix_mask, mask, g_rows=g_row, candidate_mask=( mask.astype(bool) & mask[t].astype(bool) & ( jnp.ones_like(mask, dtype=bool) if picked is None else ~picked.astype(bool) ) ), sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) return candidate, query def _teacher_logits( self, h: Float[Array, "n d_in"], edge: Float[Array, "n n d_edge"], perm: Int[Array, "n"], mask: Int[Array, "n"] | Array, *, global_feat: Float[Array, "d_global"] | None = None, tau: float | Float[Array, ""] = 1.0, real_mask: Int[Array, "n"] | Array | None = None, first_orbit_ids: QuotientCarrier, ) -> Float[Array, "n n"]: n = h.shape[0] if n > self.max_n: raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") dtype = h.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) first_active_idx = self._first_active_index(mask) node_state, _node_mean = self._prepare_nodes(h.astype(dtype), mask) global_state = self._project_global( global_feat, dtype, structural_mask=jnp.any(mask.astype(bool)), ) base = self._teacher_candidate_states( node_state, global_state, edge, perm, mask, real_mask=real_mask, ) seq = base[idx, perm, :] tree_candidate_state, hidden = self._tree_enrich_teacher( base, seq, edge, perm, mask, g=global_state[0], query_global=global_state[1], ) candidate_state = self._apply_heavy_teacher( tree_candidate_state, hidden, edge, perm, mask, ) neg = jnp.asarray(-1.0e30, dtype=dtype) pos_of_node = jnp.zeros((n,), dtype=jnp.int32).at[perm].set(idx) valid = mask_bool[None, :] & (pos_of_node[None, :] >= idx[:, None]) first_choice_mask = self._learned_first_choice_mask(mask, real_mask) valid = jnp.where( (idx == first_active_idx)[:, None], first_choice_mask[None, :], valid, ) pointer_structural_mask = mask_bool[:, None] & valid raw = self._pointer_raw( hidden, candidate_state, structural_mask=pointer_structural_mask, ) raw = raw / jnp.asarray(tau, dtype=dtype) identity = jnp.where( idx[None, :] == idx[:, None], jnp.asarray(0.0, dtype=dtype), neg, ) pointer = jnp.where(valid, raw, neg) active_scores = pointer active_scores = jax.vmap( lambda row_i, row: self._apply_quotient_logits( row, first_orbit_ids, row > (neg * jnp.asarray(0.5, dtype=dtype)), mask, perm, row_i, ) )(idx, active_scores) return jnp.where(mask_bool[:, None], active_scores, identity) def _decode( self, h: Float[Array, "n d_in"], edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, *, tau: float | Float[Array, ""], key: PRNGKeyArray, real_mask: Int[Array, "n"] | Array | None = None, first_orbit_ids: QuotientCarrier, router_static, ): n = h.shape[0] if n > self.max_n: raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") dtype = h.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) rm_bool = (real_mask if real_mask is not None else mask).astype(bool) first_active = self._first_active_index(mask) neg = jnp.asarray(-1.0e30, dtype=dtype) node_state = (router_static.node_input, router_static.node_projected) global_state = (router_static.global_input, router_static.global_projected) suffix_raw0 = router_static.initial_suffix prefix_raw0 = jnp.zeros_like(suffix_raw0) virt_count0 = jnp.zeros((), dtype=dtype) prefix_order_raw0 = jnp.zeros((n, n, self.d_model), dtype=dtype) virt_prefix_order_raw0 = jnp.zeros((n, self.d_model), dtype=dtype) order_decay = router_static.order_decay virt_decay = router_static.virtual_decay pair_msg = router_static.tree_pair_messages cross_biases, suffix_biases = self._unpack_heavy_static_bias_tables( router_static.static_bias_tables ) noise = jax.random.gumbel(key, (n, n), dtype=dtype) perm0 = idx picked0 = jnp.zeros((n,), dtype=bool) prefix_ids0 = jnp.zeros((n,), dtype=jnp.int32) k_cache0 = jnp.zeros( (0, n, self.n_heads_kernel, self.d_head), dtype=dtype, ) v_cache0 = jnp.zeros_like(k_cache0) hidden_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) base_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) last_hidden0 = jnp.zeros((self.d_model,), dtype=dtype) def body(carry, xs): ( perm, picked, prefix_raw, prefix_order_raw, suffix_raw, virt_prefix_order_raw, virt_count, prefix_ids, k_cache, v_cache, hidden_cache, base_cache, last_hidden, ) = carry t, noise_t = xs pref_msg_buf = prefix_order_raw virt_msg_buf = virt_prefix_order_raw append_step = mask_bool[t] first_step = append_step & (t == first_active) tri_t = ((idx < t) & mask_bool).astype(dtype) decay_t = order_decay[t] prefix_order_row = jnp.einsum("s,sd,sid->id", tri_t, decay_t, pref_msg_buf) vdecay_t = virt_decay[t] virt_order_row = jnp.einsum("s,sd,sd->d", tri_t, vdecay_t, virt_msg_buf) base = self._candidate_states_from_summaries( node_state, global_state, prefix_raw, prefix_order_row, suffix_raw, t, edge, mask, prefix_ids, virt_prefix_order_raw=virt_order_row, virt_count=virt_count, real_mask=real_mask, ) candidate_state, pointer_hidden = self._apply_tree_prefix_step( base, base_cache, prefix_ids, mask, t, pair_msg, g=global_state[0], query_global=global_state[1], picked=picked, ) candidate_state = self._apply_heavy_step( candidate_state, hidden_cache, edge, prefix_ids, picked, mask, t, cross_biases=cross_biases, suffix_biases=suffix_biases, ) active_logits = self._pointer_logits( pointer_hidden, candidate_state, picked, self._step_choice_mask(first_step, mask, real_mask), tau, ) active_logits = self._apply_quotient_logits( active_logits, first_orbit_ids, active_logits > (neg * jnp.asarray(0.5, dtype=dtype)), mask, prefix_ids, t, ) identity_logits = jnp.where(idx == t, jnp.asarray(0.0, dtype=dtype), neg) logits = jnp.where(append_step, active_logits, identity_logits) select_scores = logits + noise_t sampled = jnp.argmax(select_scores).astype(jnp.int32) chosen = jnp.where(append_step, sampled, t) base_chosen = base[chosen] token_in = jnp.where( append_step, base_chosen, jnp.zeros((self.d_model,), dtype=dtype), ) token, k_new, v_new = self._append_token( token_in, chosen, t, prefix_ids, k_cache, v_cache, edge, mask, ) k_cache = k_new v_cache = v_new hidden_cache = hidden_cache.at[t].set(pointer_hidden) base_cache = base_cache.at[t].set( jnp.where(append_step, base_chosen, jnp.zeros_like(base_chosen)) ) last_hidden = pointer_hidden pref_update = router_static.prefix_edge_messages[chosen] suff_update = router_static.suffix_edge_messages[chosen] update_mask = append_step.astype(dtype) prefix_raw = prefix_raw + update_mask * pref_update prefix_order_raw = pref_msg_buf.at[t].set(update_mask * pref_update) suffix_raw = suffix_raw - update_mask * suff_update virt_update = update_mask * (~rm_bool[chosen]).astype(dtype) virt_prefix_order_raw = virt_msg_buf.at[t].set( virt_update * self.virt_emb[0], ) virt_count = virt_count + virt_update prefix_ids = prefix_ids.at[t].set(chosen) picked = picked.at[chosen].set(jnp.where(append_step, True, picked[chosen])) perm = perm.at[t].set(chosen) return ( perm, picked, prefix_raw, prefix_order_raw, suffix_raw, virt_prefix_order_raw, virt_count, prefix_ids, k_cache, v_cache, hidden_cache, base_cache, last_hidden, ), None init = ( perm0, picked0, prefix_raw0, prefix_order_raw0, suffix_raw0, virt_prefix_order_raw0, virt_count0, prefix_ids0, k_cache0, v_cache0, hidden_cache0, base_cache0, last_hidden0, ) final, _ = jax.lax.scan(body, init, (idx, noise)) ( perm, _picked, _prefix_raw, _prefix_order_raw, _suffix_raw, _virt_po, _virt_cnt, _prefix_ids, _k, _v, _hidden_cache, _base_cache, _hidden, ) = final return perm def _decode_greedy_compact( self, h: Float[Array, "n d_in"], edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, *, global_feat: Float[Array, "d_global"], tau: float | Float[Array, ""], real_mask: Int[Array, "n"] | Array, sequence_mesh, pair_tile_size: int, row_permute_fn, ): n = h.shape[0] if n > self.max_n: raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") if int(pair_tile_size) < 1: raise ValueError("pair_tile_size must be positive") sequence_axis_name = "seq" dtype = h.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) rm_bool = real_mask.astype(bool) first_active = self._first_active_index(mask) neg = jnp.asarray(-1.0e30, dtype=dtype) edge_transpose = jnp.swapaxes(edge, 0, 1) from jax.sharding import NamedSharding, PartitionSpec as P edge_transpose = jax.lax.with_sharding_constraint( edge_transpose, NamedSharding(sequence_mesh, P(sequence_axis_name, None, None)), ) node_state, _node_mean = self._prepare_nodes(h.astype(dtype), mask) global_state = self._project_global( global_feat, dtype, structural_mask=jnp.any(mask_bool), ) ( prefix_raw0, _prefix_order_raw0_unused, suffix_raw0, _virt_po_unused, virt_count0, ) = self._initial_summaries_streamed( edge, edge_transpose, mask, dtype, pair_tile_size=pair_tile_size, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) pair_msg = None cross_biases, suffix_biases = self._heavy_biases_tiled( edge, edge_transpose, pair_tile_size=int(pair_tile_size), sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) frontier_depth = max(1, _route_next_pow2(n).bit_length() - 1) prefix_order_frontier0 = jnp.zeros( (frontier_depth, n, self.d_model), dtype=dtype, ) virt_order_frontier0 = jnp.zeros( (frontier_depth, self.d_model), dtype=dtype, ) source_edge_frontier0 = jnp.zeros( (frontier_depth, n, self.d_model), dtype=dtype, ) source_edge_counts0 = jnp.zeros((frontier_depth,), dtype=dtype) tree_state0 = self._incremental_tree_state(n, dtype) perm0 = idx picked0 = jnp.zeros((n,), dtype=bool) prefix_ids0 = jnp.zeros((n,), dtype=jnp.int32) k_cache0 = jnp.zeros( (0, n, self.n_heads_kernel, self.d_head), dtype=dtype, ) v_cache0 = jnp.zeros_like(k_cache0) hidden_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) base_cache0 = jnp.zeros((n, self.d_model), dtype=dtype) logp0 = jnp.asarray(0.0, dtype=jnp.float32) seq = sequence_axis_name def _seq_sharding(*axes): return NamedSharding(sequence_mesh, P(*axes)) node_state = tuple( jax.lax.with_sharding_constraint(x, _seq_sharding(seq, None)) for x in node_state ) prefix_raw0 = jax.lax.with_sharding_constraint( prefix_raw0, _seq_sharding(seq, None), ) suffix_raw0 = jax.lax.with_sharding_constraint( suffix_raw0, _seq_sharding(seq, None), ) prefix_order_frontier0 = jax.lax.with_sharding_constraint( prefix_order_frontier0, _seq_sharding(None, seq, None), ) source_edge_frontier0 = jax.lax.with_sharding_constraint( source_edge_frontier0, _seq_sharding(None, seq, None), ) hidden_cache0 = jax.lax.with_sharding_constraint( hidden_cache0, _seq_sharding(seq, None), ) base_cache0 = jax.lax.with_sharding_constraint( base_cache0, _seq_sharding(seq, None), ) k_cache0 = jax.lax.with_sharding_constraint( k_cache0, _seq_sharding(None, seq, None, None), ) v_cache0 = jax.lax.with_sharding_constraint( v_cache0, _seq_sharding(None, seq, None, None), ) edge_transpose = jax.lax.with_sharding_constraint( edge_transpose, _seq_sharding(seq, None, None), ) tree_nodes_raw0, tree_nodes_post0, tree_levels0 = tree_state0 tree_nodes_raw0 = jax.lax.with_sharding_constraint( tree_nodes_raw0, _seq_sharding(None, seq, None), ) tree_nodes_post0 = jax.lax.with_sharding_constraint( tree_nodes_post0, _seq_sharding(None, seq, None), ) lanes = int(sequence_mesh.shape[seq]) constrained_levels = [] for edge_pre0, edge_post0, b_cache0 in tree_levels0: shard_rows = ( seq if edge_pre0.shape[0] >= lanes and edge_pre0.shape[0] % lanes == 0 else None ) constrained_levels.append( ( jax.lax.with_sharding_constraint( edge_pre0, _seq_sharding(shard_rows, None, None), ), jax.lax.with_sharding_constraint( edge_post0, _seq_sharding(shard_rows, None, None), ), jax.lax.with_sharding_constraint( b_cache0, _seq_sharding(shard_rows, None, None), ), ) ) tree_state0 = ( tree_nodes_raw0, tree_nodes_post0, tuple(constrained_levels), ) def body(carry, t): ( perm, picked, prefix_raw, prefix_order_frontier, suffix_raw, virt_order_frontier, virt_count, source_edge_frontier, source_edge_counts, tree_state, prefix_ids, k_cache, v_cache, hidden_cache, base_cache, total_logp, ) = carry append_step = mask_bool[t] first_step = append_step & (t == first_active) predict_step = append_step & (t != first_active) prefix_order_row = _dyadic_lca_frontier_sum( prefix_order_frontier, t, self.order_decay_w[0], self.order_decay_b[0], ) virt_order_row = _dyadic_lca_frontier_sum( virt_order_frontier, t, self.virt_decay_w[0], self.virt_decay_b[0], ) base = self._candidate_states_from_summaries( node_state, global_state, prefix_raw, prefix_order_row, suffix_raw, t, edge, mask, prefix_ids, virt_prefix_order_raw=virt_order_row, virt_count=virt_count, real_mask=real_mask, ) candidate_state, pointer_hidden = self._apply_tree_prefix_step( base, base_cache, prefix_ids, mask, t, pair_msg, g=global_state[0], query_global=global_state[1], picked=picked, source_edge_frontier=source_edge_frontier, source_edge_counts=source_edge_counts, raw_edge_for_pair_messages=edge, raw_edge_transpose=edge_transpose, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, row_permute_fn=row_permute_fn, incremental_tree_state=tree_state, ) candidate_state = self._apply_heavy_step( candidate_state, hidden_cache, edge, prefix_ids, picked, mask, t, cross_biases=cross_biases, suffix_biases=suffix_biases, edge_transpose=edge_transpose, sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) active_logits = self._pointer_logits( pointer_hidden, candidate_state, picked, self._step_choice_mask(first_step, mask, real_mask), tau, ) identity_logits = jnp.where( idx == t, jnp.asarray(0.0, dtype=dtype), neg, ) logits = jnp.where(append_step, active_logits, identity_logits) sampled = jnp.argmax(logits).astype(jnp.int32) chosen = jnp.where(append_step, sampled, t) log_probs = jax.nn.log_softmax(logits.astype(jnp.float32), axis=-1) score_step = self._score_step_for_logp(first_step, predict_step) total_logp = total_logp + jnp.where( score_step, log_probs[chosen], 0.0, ) base_chosen = base[chosen] token_in = jnp.where( append_step, base_chosen, jnp.zeros((self.d_model,), dtype=dtype), ) _token, k_cache, v_cache = self._append_token( token_in, chosen, t, prefix_ids, k_cache, v_cache, edge, mask, ) hidden_cache = hidden_cache.at[t].set(pointer_hidden) base_cache = base_cache.at[t].set( jnp.where(append_step, base_chosen, jnp.zeros_like(base_chosen)) ) chosen_edge_pair = jnp.concatenate( [edge[:, chosen, :], edge_transpose[:, chosen, :]], axis=-1, ) pref_update = self._message_mlp(chosen_edge_pair, prefix=True) suff_update = self._message_mlp(chosen_edge_pair, prefix=False) update_mask = append_step.astype(dtype) prefix_raw = prefix_raw + update_mask * pref_update prefix_order_frontier = _dyadic_frontier_add( prefix_order_frontier, update_mask * pref_update, t, ) suffix_raw = suffix_raw - update_mask * suff_update virt_update = update_mask * (~rm_bool[chosen]).astype(dtype) virt_order_frontier = _dyadic_frontier_add( virt_order_frontier, virt_update * self.virt_emb[0], t, ) virt_count = virt_count + virt_update source_edge_update = self._tree_pair_message_row( edge, chosen, mask, edge_transpose=edge_transpose, ) source_edge_frontier = _dyadic_frontier_add( source_edge_frontier, update_mask * source_edge_update, t, ) source_edge_counts = _dyadic_frontier_add( source_edge_counts, update_mask, t, ) prefix_ids = prefix_ids.at[t].set(chosen) tree_state = self._incremental_tree_append( tree_state, jnp.where( append_step, base_chosen, jnp.zeros_like(base_chosen), ), chosen, t, prefix_ids, edge, mask, edge_transpose=edge_transpose, g=global_state[0], sequence_axis_name=sequence_axis_name, sequence_mesh=sequence_mesh, ) picked = picked.at[chosen].set(jnp.where(append_step, True, picked[chosen])) perm = perm.at[t].set(chosen) next_carry = ( perm, picked, prefix_raw, prefix_order_frontier, suffix_raw, virt_order_frontier, virt_count, source_edge_frontier, source_edge_counts, tree_state, prefix_ids, k_cache, v_cache, hidden_cache, base_cache, total_logp, ) return next_carry, None init = ( perm0, picked0, prefix_raw0, prefix_order_frontier0, suffix_raw0, virt_order_frontier0, virt_count0, source_edge_frontier0, source_edge_counts0, tree_state0, prefix_ids0, k_cache0, v_cache0, hidden_cache0, base_cache0, logp0, ) final, _ = jax.lax.scan(body, init, idx) perm = final[0] logp = final[-1] return perm, logp def beam_search( self, h: Float[Array, "n d_in"], edge: Float[Array, "n n d_edge"], mask: Int[Array, "n"] | Array, *, global_feat: Float[Array, "d_global"] | None = None, tau: float | Float[Array, ""] = 1.0, beam_width: int = 4, real_mask: Int[Array, "n"] | Array | None = None, first_orbit_ids: QuotientCarrier, router_static=None, distributed_axis_name: str | None = None, distributed_lanes: int | None = None, ): B = int(beam_width) if B < 1: raise ValueError("beam_width must be >= 1") if distributed_axis_name is not None: lanes = int(distributed_lanes if distributed_lanes is not None else 8) if lanes < 1 or B % lanes: raise ValueError( f"distributed beam requires beam_width divisible by the " f"lane count; got beam_width={B}, lanes={lanes}" ) if router_static is None: raise ValueError("distributed audit beam requires RouterStatic") n = h.shape[0] if n > self.max_n: raise ValueError(f"route pointer saw N={n} > max_n={self.max_n}") dtype = h.dtype idx = jnp.arange(n, dtype=jnp.int32) mask_bool = mask.astype(bool) rm_bool = (real_mask if real_mask is not None else mask).astype(bool) first_active = self._first_active_index(mask) neg = jnp.asarray(-1.0e30, dtype=dtype) if router_static is None: node_state, _node_mean = self._prepare_nodes(h.astype(dtype), mask) global_state = self._project_global( global_feat, dtype, structural_mask=jnp.any(mask.astype(bool)), ) ( prefix_raw0, _prefix_order_raw0_unused, suffix_raw0, _virt_po_unused, virt_count0, ) = self._initial_summaries(edge, mask, dtype) else: node_state = (router_static.node_input, router_static.node_projected) global_state = (router_static.global_input, router_static.global_projected) suffix_raw0 = router_static.initial_suffix prefix_raw0 = jnp.zeros_like(suffix_raw0) virt_count0 = jnp.zeros((), dtype=dtype) prefix_order_raw0 = jnp.zeros((n, n, self.d_model), dtype=dtype) virt_prefix_order_raw0 = jnp.zeros((n, self.d_model), dtype=dtype) if router_static is None: order_decay = lca_gaussian_decay( idx, idx, self.order_decay_w[0], self.order_decay_b[0], ) virt_decay = lca_gaussian_decay( idx, idx, self.virt_decay_w[0], self.virt_decay_b[0], ) pair_msg = self._tree_pair_messages(edge, mask) else: order_decay = router_static.order_decay virt_decay = router_static.virtual_decay pair_msg = router_static.tree_pair_messages if router_static is None: cross_biases = self._heavy_cross_biases(edge) suffix_biases = self._heavy_suffix_biases(edge) else: cross_biases, suffix_biases = self._unpack_heavy_static_bias_tables( router_static.static_bias_tables ) def repeat(x): return jnp.broadcast_to(x, (B,) + x.shape) perm0 = repeat(idx) picked0 = jnp.zeros((B, n), dtype=bool) prefix_ids0 = jnp.zeros((B, n), dtype=jnp.int32) k_cache0 = jnp.zeros( (B, 0, n, self.n_heads_kernel, self.d_head), dtype=dtype, ) v_cache0 = jnp.zeros_like(k_cache0) hidden_cache0 = jnp.zeros((B, n, self.d_model), dtype=dtype) base_cache0 = jnp.zeros((B, n, self.d_model), dtype=dtype) last_hidden0 = jnp.zeros((B, self.d_model), dtype=dtype) logp0 = ( jnp.full((B,), jnp.asarray(-1e9, dtype=jnp.float32), dtype=jnp.float32) .at[0] .set(0.0) ) beam_ids = jnp.arange(B, dtype=jnp.int32) rows = jnp.arange(B, dtype=jnp.int32) def body(carry, t): ( perm, picked, prefix_raw, prefix_order_raw, suffix_raw, virt_prefix_order_raw, virt_count, prefix_ids, k_cache, v_cache, hidden_cache, base_cache, last_hidden, total_logp, ) = carry append_step = mask_bool[t] first_step = append_step & (t == first_active) predict_step = append_step & (t != first_active) def states_one( pr, por_buf, sr, vpo_buf, vcnt, pids, pk, bc, hcache, hidden ): tri_t = ((idx < t) & mask_bool).astype(dtype) decay_t = order_decay[t] por = jnp.einsum("s,sd,sid->id", tri_t, decay_t, por_buf) vdecay_t = virt_decay[t] virt_order_row = jnp.einsum("s,sd,sd->d", tri_t, vdecay_t, vpo_buf) base = self._candidate_states_from_summaries( node_state, global_state, pr, por, sr, t, edge, mask, pids, virt_prefix_order_raw=virt_order_row, virt_count=vcnt, real_mask=real_mask, ) candidate_state, pointer_hidden = self._apply_tree_prefix_step( base, bc, pids, mask, t, pair_msg, g=global_state[0], query_global=global_state[1], picked=pk, ) candidate_state = self._apply_heavy_step( candidate_state, hcache, edge, pids, pk, mask, t, cross_biases=cross_biases, suffix_biases=suffix_biases, ) active_logits = self._pointer_logits( pointer_hidden, candidate_state, pk, self._step_choice_mask(first_step, mask, real_mask), tau, ) active_logits = self._apply_quotient_logits( active_logits, first_orbit_ids, active_logits > (neg * jnp.asarray(0.5, dtype=dtype)), mask, pids, t, ) identity_logits = jnp.where( idx == t, jnp.asarray(0.0, dtype=dtype), neg ) logits = jnp.where(append_step, active_logits, identity_logits) return logits, candidate_state, base, pointer_hidden if distributed_axis_name is None: parent_rows = rows else: lane = jax.lax.axis_index(distributed_axis_name) _per_lane = B // lanes parent_rows = lane * _per_lane + jnp.arange(_per_lane, dtype=jnp.int32) ( logits_local, _candidate_state_local, base_state_local, query_state_local, ) = jax.vmap(states_one)( prefix_raw[parent_rows], prefix_order_raw[parent_rows], suffix_raw[parent_rows], virt_prefix_order_raw[parent_rows], virt_count[parent_rows], prefix_ids[parent_rows], picked[parent_rows], base_cache[parent_rows], hidden_cache[parent_rows], last_hidden[parent_rows], ) score_step = self._score_step_for_logp(first_step, predict_step) neg_f32 = jnp.asarray(-1e9, dtype=jnp.float32) def expansion_for(logits_arg, parent_total): log_probs_arg = jax.nn.log_softmax( logits_arg.astype(jnp.float32), axis=-1 ) step_logp_arg = jnp.where( score_step, log_probs_arg, jnp.zeros_like(log_probs_arg) ) expansion_arg = parent_total[:, None] + step_logp_arg forced_scores = jnp.where( idx[None, :] == t.astype(jnp.int32), expansion_arg, neg_f32, ) return jnp.where(~append_step, forced_scores, expansion_arg) expansion_local = expansion_for(logits_local, total_logp[parent_rows]) if distributed_axis_name is None: expansion_scores = expansion_local base_state = base_state_local query_state = query_state_local else: _pl = B // lanes base_shape = base_state_local.shape query_shape = query_state_local.shape payload_parts = [ expansion_local.reshape((_pl, -1)), ] payload_parts.extend( [ base_state_local.astype(jnp.float32).reshape((_pl, -1)), query_state_local.astype(jnp.float32).reshape((_pl, -1)), ] ) payload = jnp.concatenate(payload_parts, axis=-1) payload = jax.lax.all_gather( payload, distributed_axis_name, axis=0, tiled=True, ) cursor = 0 expansion_scores = payload[:, cursor : cursor + n] cursor += n base_size = n * base_shape[-1] base_state = ( payload[:, cursor : cursor + base_size] .reshape((B, n, base_shape[-1])) .astype(dtype) ) cursor += base_size query_state = payload[:, cursor : cursor + query_shape[-1]].astype( dtype ) rank_scores = expansion_scores identity_distance = jnp.abs(idx - t.astype(jnp.int32)).astype(jnp.float32) rank_scores = rank_scores - identity_distance[None, :] * 1.0e-6 rank_scores = rank_scores - beam_ids[:, None].astype(jnp.float32) * 1.0e-9 _rank_top, flat = jax.lax.top_k(rank_scores.reshape((-1,)), B) parent = (flat // n).astype(jnp.int32) chosen = (flat % n).astype(jnp.int32) total_logp = expansion_scores.reshape((-1,))[flat] perm = perm[parent] picked = picked[parent] prefix_raw = prefix_raw[parent] prefix_order_raw = prefix_order_raw[parent] suffix_raw = suffix_raw[parent] virt_prefix_order_raw = virt_prefix_order_raw[parent] virt_count = virt_count[parent] prefix_ids = prefix_ids[parent] k_cache = k_cache[parent] v_cache = v_cache[parent] hidden_cache = hidden_cache[parent] base_cache = base_cache[parent] base_chosen = base_state[parent, chosen] query_chosen = query_state[parent] token_in = jnp.where( append_step, base_chosen, jnp.zeros_like(base_chosen), ) token, k_cache, v_cache = jax.vmap( lambda token_b, chosen_b, prefix_ids_b, k_b, v_b: self._append_token( token_b, chosen_b, t, prefix_ids_b, k_b, v_b, edge, mask, ) )(token_in, chosen, prefix_ids, k_cache, v_cache) hidden_cache = hidden_cache.at[:, t, :].set(query_chosen) base_cache = base_cache.at[:, t, :].set( append_step.astype(dtype) * base_chosen ) last_hidden = query_chosen if router_static is None: chosen_edge_pair = jax.vmap( lambda chosen_b: self._edge_pair_for_source(edge, chosen_b) )(chosen) pref_update = jax.vmap( lambda pair: self._message_mlp(pair, prefix=True) )(chosen_edge_pair) suff_update = jax.vmap( lambda pair: self._message_mlp(pair, prefix=False) )(chosen_edge_pair) else: pref_update = router_static.prefix_edge_messages[chosen] suff_update = router_static.suffix_edge_messages[chosen] update_mask = append_step.astype(dtype) prefix_raw = prefix_raw + update_mask * pref_update prefix_order_raw = prefix_order_raw.at[:, t].set(update_mask * pref_update) suffix_raw = suffix_raw - update_mask * suff_update virt_update = update_mask * (~rm_bool[chosen]).astype(dtype) virt_prefix_order_raw = virt_prefix_order_raw.at[:, t].set( virt_update[:, None] * self.virt_emb[0][None, :], ) virt_count = virt_count + virt_update prefix_ids = prefix_ids.at[:, t].set(chosen) old_picked = picked[rows, chosen] picked = picked.at[rows, chosen].set( jnp.where(append_step, True, old_picked) ) perm = perm.at[:, t].set(chosen) return ( perm, picked, prefix_raw, prefix_order_raw, suffix_raw, virt_prefix_order_raw, virt_count, prefix_ids, k_cache, v_cache, hidden_cache, base_cache, last_hidden, total_logp, ), None init = ( perm0, picked0, repeat(prefix_raw0), repeat(prefix_order_raw0), repeat(suffix_raw0), repeat(virt_prefix_order_raw0), repeat(virt_count0), prefix_ids0, k_cache0, v_cache0, hidden_cache0, base_cache0, last_hidden0, logp0, ) final, _ = jax.lax.scan(body, init, idx) ( perm, _picked, _prefix_raw, _prefix_order_raw, _suffix_raw, _virt_po, _virt_cnt, _prefix_ids, _k, _v, _hidden_cache, _base_cache, _hidden, logp, ) = final return perm, logp __all__ = ["TreePrefixPointerMHSEA"]