KitsuVp commited on
Commit
b7c5fa6
·
verified ·
1 Parent(s): 774440e

Update modeling_neollm.py

Browse files
Files changed (1) hide show
  1. modeling_neollm.py +756 -30
modeling_neollm.py CHANGED
@@ -239,6 +239,10 @@ class NeoLLMBaseModelOutputWithPast(BaseModelOutputWithPast):
239
  jtokm_aux_stats: Optional[
240
  Tuple[Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...]
241
  ] = None
 
 
 
 
242
 
243
  logger = logging.get_logger(__name__)
244
  _NEOLLM_FA_USE_TOP_LEFT_MASK = flash_attn_supports_top_left_mask()
@@ -403,6 +407,394 @@ class EmbeddingWithMultipliers(nn.Module):
403
  # ==================== LEVIATHAN CONTINUOUS TOKEN GENERATOR ====================
404
 
405
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
406
  class LeviathanGenerator(nn.Module):
407
  """
408
  Learned Embedding Vectorization (LEV) layer - Leviathan input embedding
@@ -451,6 +843,11 @@ class LeviathanGenerator(nn.Module):
451
  # is unavailable; callers can also disable the optional route per
452
  # instance without changing the model/state-dict contract.
453
  self.use_leviathan_triton = bool(_LEV_KERNEL_AVAILABLE)
 
 
 
 
 
454
  vocab_size = config.vocab_size
455
  hidden_size = config.hidden_size
456
  self.d_seed = config.generator_d_seed
@@ -639,6 +1036,18 @@ class LeviathanGenerator(nn.Module):
639
  prod_sign = 1.0 - 2.0 * (num_neg % 2).float() # [N, krank]
640
  return prod_sign * torch.exp(log_mag) # [N, krank]
641
 
 
 
 
 
 
 
 
 
 
 
 
 
642
  # Keep this method traceable. The compiler-safe LEV entry point registers
643
  # an opaque custom op tagged cudagraph_unsafe: Dynamo keeps the LEV call
644
  # in the compiled graph while Inductor excludes only the unsafe Triton
@@ -651,6 +1060,7 @@ class LeviathanGenerator(nn.Module):
651
  meap_mask_embedding: Optional[torch.Tensor] = None,
652
  meap_mask_token_id: Optional[int] = None,
653
  return_jtok_geometry: bool = False,
 
654
  ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
655
  """
656
  Generate embeddings from discrete token indices.
@@ -665,6 +1075,11 @@ class LeviathanGenerator(nn.Module):
665
  where ``z_tilde`` is the shared differentiable Leviathan/JTok
666
  coordinate with shape ``[*token_ids.shape, d_seed]``.
667
  """
 
 
 
 
 
668
  has_meap = (
669
  meap_mask_embedding is not None or meap_mask_token_id is not None
670
  )
@@ -704,7 +1119,7 @@ class LeviathanGenerator(nn.Module):
704
  )
705
  require_leviathan_triton = bool(
706
  getattr(self.config, "_require_leviathan_triton", False)
707
- )
708
  if require_leviathan_triton and not leviathan_kernel_ready:
709
  raise RuntimeError(
710
  "This run requires the Triton Leviathan kernel, but the "
@@ -749,24 +1164,38 @@ class LeviathanGenerator(nn.Module):
749
  "but the installed CCE package lacks the geometry "
750
  "adapter."
751
  )
752
- embedding, seed = _leviathan_embedding_with_seed_compiler_safe(
753
  token_ids,
754
  params,
755
  self.config,
756
  self.knot_grid,
757
  mask_embedding=meap_mask_embedding,
758
  mask_token_id=meap_mask_token_id,
 
759
  )
 
 
 
 
 
 
760
  embedding = embedding.reshape(*token_ids.shape, self.hidden_size)
761
  return embedding, self._jtok_coordinate_from_seed(seed)
762
- return _leviathan_embedding_compiler_safe(
763
  token_ids,
764
  params,
765
  self.config,
766
  self.knot_grid,
767
  mask_embedding=meap_mask_embedding,
768
  mask_token_id=meap_mask_token_id,
 
769
  )
 
 
 
 
 
 
770
  except (TypeError, ValueError, AttributeError):
771
  # Only validation/configuration failures use the reference
772
  # implementation. A CUDA launch error must propagate: doing
@@ -870,6 +1299,11 @@ class LeviathanJTok(nn.Module):
870
  super().__init__()
871
  self.layer_idx = int(layer_idx)
872
  self.use_mixture = bool(getattr(config, "use_jtokm", False))
 
 
 
 
 
873
  self.jtok_kernel_backend = str(
874
  getattr(config, "jtok_kernel_backend", "torch")
875
  ).strip().lower()
@@ -877,6 +1311,8 @@ class LeviathanJTok(nn.Module):
877
  raise ValueError(
878
  "`jtok_kernel_backend` must be either `torch` or `triton`."
879
  )
 
 
880
  if self.jtok_kernel_backend == "triton" and not _JTOK_KERNEL_AVAILABLE:
881
  raise RuntimeError(
882
  "`jtok_kernel_backend='triton'` was requested, but the "
@@ -1104,6 +1540,7 @@ class LeviathanJTok(nn.Module):
1104
  Optional[Tuple[torch.Tensor, torch.Tensor, torch.Tensor]],
1105
  ]:
1106
  """Apply JTok/JTok-M and optionally return training-only router stats."""
 
1107
  if self.jtok_kernel_backend == "triton":
1108
  # Dispatch only: this branch does not change compiler mode,
1109
  # CUDA-graph policy, training state, optimizer behavior, or any
@@ -1116,7 +1553,7 @@ class LeviathanJTok(nn.Module):
1116
  "JTok Triton adapter became unavailable after module "
1117
  "construction; refusing an implicit Torch fallback."
1118
  )
1119
- return _apply_neollm_jtok(
1120
  self,
1121
  delta_m,
1122
  z_tilde,
@@ -1124,7 +1561,22 @@ class LeviathanJTok(nn.Module):
1124
  valid_mask=valid_mask,
1125
  compute_aux=bool(compute_aux and self.use_mixture),
1126
  backend="triton",
1127
- )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1128
 
1129
  orig_shape = delta_m.shape
1130
  N = delta_m.numel() // self.hidden_size
@@ -1286,7 +1738,11 @@ class StackMemory(nn.Module):
1286
  self.down_proj = nn.Linear(config.hidden_size, config.stack_d_model)
1287
  self.up_proj = nn.Linear(config.stack_d_model, config.hidden_size)
1288
  self.action_head = nn.Linear(config.stack_d_model, 3 * self.num_mem_heads)
1289
- self.gate_proj = nn.Linear(self.stack_dim, 1)
 
 
 
 
1290
  self.res_weight = nn.Parameter(torch.ones(1))
1291
 
1292
  self.cache_size = getattr(config, "stack_memory_cache_size", 2048)
@@ -1311,24 +1767,37 @@ class StackMemory(nn.Module):
1311
  [mask[:, :, :, 1:], torch.zeros_like(mask[:, :, :, :1])], dim=3
1312
  )
1313
 
1314
- action_weights = actions.unsqueeze(-1).unsqueeze(-1)
 
 
 
 
1315
  stacks = torch.stack([push_stack, pop_stack, stack], dim=3)
1316
  masks = torch.stack([push_mask, pop_mask, mask], dim=3)
1317
 
1318
- new_stack = (stacks * action_weights).sum(dim=3)
1319
- new_mask = (masks * action_weights.squeeze(-1)).sum(dim=3)
1320
  return new_stack, new_mask
1321
 
1322
- def forward(self, hidden_states, stack, mask):
1323
  batch_size, seq_len, _ = hidden_states.shape
1324
 
 
 
 
 
1325
  # OLMo stack_v3 value/action path: project the complete hidden state into
1326
  # the total low-rank stack width, then split that representation by head.
1327
  new_hidden_states = self.down_proj(hidden_states)
1328
 
1329
  action_logits = self.action_head(new_hidden_states) / math.sqrt(self.stack_dim)
 
 
 
1330
  actions = F.softmax(
1331
- action_logits.reshape(batch_size, seq_len, self.num_mem_heads, 3),
 
 
1332
  dim=-1,
1333
  )
1334
 
@@ -1338,21 +1807,99 @@ class StackMemory(nn.Module):
1338
 
1339
  new_stack, new_mask = self._vectorized_update(stack, mask, actions, k_values)
1340
 
1341
- gate_scores = self.gate_proj(new_stack).squeeze(-1)
1342
- # Keep NeoLLM's requested softened source mask exactly at -80.
1343
- gate_weights = F.softmax(gate_scores + (1 - new_mask) * -80.0, dim=-1)
 
 
 
 
 
 
 
 
 
 
1344
 
1345
- memory_output = (new_stack * gate_weights.unsqueeze(-1)).sum(dim=3)
 
1346
  memory_output = memory_output.reshape(batch_size, seq_len, -1)
1347
  memory_output = self.up_proj(memory_output)
1348
 
1349
  output = memory_output * self.res_weight + hidden_states
1350
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1351
  if self.training and self.enable_cache:
1352
  self._update_cache(k_values.detach(), actions.detach())
1353
 
1354
  # Preserve the complete token-local stack for the next decoder layer.
1355
- return output, new_stack, new_mask
1356
 
1357
  def _update_cache(self, k_values, actions):
1358
  seq_len = k_values.shape[1]
@@ -1376,7 +1923,7 @@ class StackMemory(nn.Module):
1376
  if mask.ndim == 3:
1377
  mask = mask.unsqueeze(1)
1378
 
1379
- output, new_stack, new_mask = self.forward(
1380
  hidden_state.unsqueeze(1), stack, mask
1381
  )
1382
  return output.squeeze(1), new_stack.squeeze(1), new_mask.squeeze(1)
@@ -5487,7 +6034,13 @@ class NeoLLMDecoderLayer(GradientCheckpointingLayer):
5487
  hidden_states: torch.Tensor,
5488
  stack_memory: Optional[torch.Tensor],
5489
  stack_memory_mask: Optional[torch.Tensor],
5490
- ) -> Tuple[torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]:
 
 
 
 
 
 
5491
  """
5492
  Apply this layer's differentiable StackMemory module before attention.
5493
 
@@ -5497,13 +6050,15 @@ class NeoLLMDecoderLayer(GradientCheckpointingLayer):
5497
  attention implementation remains unchanged.
5498
  """
5499
  if (not self.use_stack_memory) or self.stack_memory is None:
5500
- return hidden_states, stack_memory, stack_memory_mask
5501
  if stack_memory is None or stack_memory_mask is None:
5502
  raise ValueError(
5503
  "StackMemory is enabled, but stack_memory/stack_memory_mask "
5504
  "were not initialized by NeoLLMModel.forward."
5505
  )
5506
- return self.stack_memory(hidden_states, stack_memory, stack_memory_mask)
 
 
5507
 
5508
  def _attn_res(
5509
  self,
@@ -5624,6 +6179,7 @@ class NeoLLMDecoderLayer(GradientCheckpointingLayer):
5624
  position_embeddings: tuple[torch.Tensor, torch.Tensor],
5625
  stack_memory: Optional[torch.Tensor] = None,
5626
  stack_memory_mask: Optional[torch.Tensor] = None,
 
5627
  attention_mask: Optional[torch.Tensor] = None,
5628
  first_layer_fan: Optional[torch.Tensor] = None,
5629
  output_attentions: Optional[bool] = False,
@@ -5647,9 +6203,15 @@ class NeoLLMDecoderLayer(GradientCheckpointingLayer):
5647
  if self.siamese_attn_input_norm is not None:
5648
  h_attn = self.siamese_attn_input_norm(h_attn)
5649
 
 
5650
  if self.use_stack_memory:
5651
- h_attn, stack_memory, stack_memory_mask = self.apply_stack_memory(
5652
- h_attn, stack_memory, stack_memory_mask
 
 
 
 
 
5653
  )
5654
  jtok_router_state = h_attn
5655
 
@@ -5706,7 +6268,7 @@ class NeoLLMDecoderLayer(GradientCheckpointingLayer):
5706
 
5707
  outputs = (x_next, y_next)
5708
  if self.use_stack_memory:
5709
- outputs += (stack_memory, stack_memory_mask)
5710
  if jtok_aux_stats is not None:
5711
  outputs += jtok_aux_stats
5712
  if output_attentions:
@@ -6637,6 +7199,22 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6637
  )
6638
  jtok_z_tilde = None
6639
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6640
  # ── Embedding stage ────────────────────────────────────────────────
6641
  meap_mask_embedding = None
6642
  meap_mask_token_id = None
@@ -6665,6 +7243,7 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6665
  meap_mask_embedding=meap_mask_embedding,
6666
  meap_mask_token_id=meap_mask_token_id,
6667
  return_jtok_geometry=True,
 
6668
  )
6669
  else:
6670
  # Keep the pre-JTok call signature for custom or
@@ -6673,15 +7252,17 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6673
  input_ids,
6674
  meap_mask_embedding=meap_mask_embedding,
6675
  meap_mask_token_id=meap_mask_token_id,
 
6676
  )
6677
  else:
6678
  if use_jtok:
6679
  generator_result = self.token_generator(
6680
  input_ids,
6681
  return_jtok_geometry=True,
 
6682
  )
6683
  else:
6684
- generator_result = self.token_generator(input_ids)
6685
  if use_jtok:
6686
  inputs_embeds, jtok_z_tilde = generator_result
6687
  else:
@@ -6820,12 +7401,17 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6820
  .to(dtype=hidden_states.dtype)
6821
  .expand(batch_size, seq_len, -1, -1, -1)
6822
  )
6823
- stack_memory_mask = self.memory_mask.detach().expand(
6824
- batch_size, seq_len, -1, -1
 
 
 
 
6825
  )
6826
  else:
6827
  stack_memory = None
6828
  stack_memory_mask = None
 
6829
 
6830
  position_embeddings = self.rotary_emb(hidden_states, position_ids)
6831
  self.first_layer_fan = (
@@ -6901,11 +7487,15 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6901
  attn_res_partial = hidden_states # start new block from current output
6902
 
6903
  if use_stack_memory and not use_siamesenorm:
6904
- hidden_states, stack_memory, stack_memory_mask = (
6905
  decoder_layer.apply_stack_memory(
6906
- hidden_states, stack_memory, stack_memory_mask
 
 
 
6907
  )
6908
  )
 
6909
 
6910
  if use_siamesenorm:
6911
  layer_outputs = decoder_layer.forward_siamesenorm(
@@ -6914,6 +7504,7 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6914
  position_embeddings=position_embeddings,
6915
  stack_memory=stack_memory,
6916
  stack_memory_mask=stack_memory_mask,
 
6917
  attention_mask=causal_mask,
6918
  first_layer_fan=self.first_layer_fan,
6919
  output_attentions=output_attentions,
@@ -6929,7 +7520,9 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6929
  if use_stack_memory:
6930
  stack_memory = layer_outputs[2]
6931
  stack_memory_mask = layer_outputs[3]
6932
- extras_start = 4
 
 
6933
  else:
6934
  extras_start = 2
6935
  else:
@@ -6990,12 +7583,20 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
6990
  if output_hidden_states:
6991
  all_hidden_states = all_hidden_states + (hidden_states,)
6992
 
 
 
 
 
 
 
6993
  if not return_dict:
6994
  base_outputs = tuple(
6995
  v
6996
  for v in [hidden_states, None, all_hidden_states, all_attentions]
6997
  if v is not None
6998
  )
 
 
6999
  if all_jtokm_aux_stats is not None:
7000
  base_outputs = base_outputs + (all_jtokm_aux_stats,)
7001
  if output_tweo_activations:
@@ -7008,6 +7609,7 @@ class NeoLLMModel(NeoLLMPreTrainedModel):
7008
  hidden_states=all_hidden_states,
7009
  attentions=all_attentions,
7010
  jtokm_aux_stats=all_jtokm_aux_stats,
 
7011
  )
7012
  if output_tweo_activations:
7013
  return outputs, tweo_activations
@@ -8840,6 +9442,20 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
8840
  self._last_nextlat_mse_loss = None
8841
  self._last_nextlat_kl_loss = None
8842
  self._last_jtokm_aux_loss = None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8843
  self._last_total_loss = None
8844
 
8845
  if config.use_token_generator:
@@ -8850,6 +9466,50 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
8850
  def get_input_embeddings(self):
8851
  return self.model.get_input_embeddings()
8852
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8853
  def set_input_embeddings(self, value):
8854
  self.model.set_input_embeddings(value)
8855
 
@@ -9025,6 +9685,7 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9025
  tweo_activations = None
9026
  hidden_states_tuple = None
9027
  jtokm_aux_stats = None
 
9028
  if isinstance(model_out, tuple):
9029
  outputs = model_out[0]
9030
  if tweo_enabled and len(model_out) >= 2:
@@ -9041,6 +9702,13 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9041
  if tuple_candidates:
9042
  hidden_states_tuple = tuple_candidates[0]
9043
  for candidate in model_out[1:]:
 
 
 
 
 
 
 
9044
  if (
9045
  isinstance(candidate, tuple)
9046
  and candidate
@@ -9048,7 +9716,6 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9048
  and len(candidate[0]) == 3
9049
  ):
9050
  jtokm_aux_stats = candidate
9051
- break
9052
  elif isinstance(outputs, tuple):
9053
  hidden_states = outputs[0]
9054
  if len(outputs) > 2:
@@ -9057,11 +9724,13 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9057
  hidden_states = outputs.last_hidden_state
9058
  hidden_states_tuple = outputs.hidden_states
9059
  jtokm_aux_stats = getattr(outputs, "jtokm_aux_stats", None)
 
9060
  else:
9061
  outputs = model_out
9062
  hidden_states = outputs.last_hidden_state
9063
  hidden_states_tuple = outputs.hidden_states
9064
  jtokm_aux_stats = getattr(outputs, "jtokm_aux_stats", None)
 
9065
 
9066
  loss = None
9067
  ntp_loss = None
@@ -9085,6 +9754,7 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9085
  nextlat_mse_loss = None
9086
  nextlat_kl_loss = None
9087
  jtokm_aux_loss = None
 
9088
  self._last_ntp_loss = None
9089
  self._last_ntp_ce_unweighted = None
9090
  self._last_mile_reweighting_delta = None
@@ -9111,6 +9781,20 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9111
  self._last_nextlat_mse_loss = None
9112
  self._last_nextlat_kl_loss = None
9113
  self._last_jtokm_aux_loss = None
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9114
  self._last_total_loss = None
9115
  if labels is not None:
9116
  if self.ntp_loss_backend == "liger":
@@ -9166,6 +9850,24 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9166
  ntp_loss = ntp_result
9167
  loss = ntp_loss
9168
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9169
  if tweo_enabled:
9170
  if not tweo_activations:
9171
  raise ValueError(
@@ -9451,6 +10153,30 @@ class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin):
9451
  self._last_jtokm_aux_loss = (
9452
  jtokm_aux_loss.detach() if jtokm_aux_loss is not None else None
9453
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9454
  self._last_total_loss = loss.detach() if loss is not None else None
9455
  logits = None
9456
  else:
 
239
  jtokm_aux_stats: Optional[
240
  Tuple[Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...]
241
  ] = None
242
+ # StackMemory diagnostics packed as [num_layers, 11]. Column 0 retains
243
+ # the differentiable action entropy used by the auxiliary StackTrans loss;
244
+ # the remaining columns are detached monitoring reductions.
245
+ stack_metrics: Optional[torch.Tensor] = None
246
 
247
  logger = logging.get_logger(__name__)
248
  _NEOLLM_FA_USE_TOP_LEFT_MASK = flash_attn_supports_top_left_mask()
 
407
  # ==================== LEVIATHAN CONTINUOUS TOKEN GENERATOR ====================
408
 
409
 
410
+ # Sampled dynamics: mathematical definitions and interpretation
411
+ # ----------------------------------------------------------
412
+ # Enable through USE_DYNAMICS_METRICS in train.py before model construction.
413
+ # train.py supplies config.dynamics_metrics_enabled and dynamics_sample_tokens;
414
+ # a standalone caller may set those same attributes before constructing a model.
415
+ # Collection requires native CUDA/Triton checkpoints, training mode and enabled
416
+ # autograd. It is detached, read-only, and never contributes to the objective.
417
+ #
418
+ # Sampling/aggregation contract:
419
+ # Flatten batch and sequence to N rows, choose S=min(N, sample_tokens) positions
420
+ # r_j=floor(j*N/S), then exclude invalid sampled rows. There is no resampling
421
+ # to replace masked rows. Explicit attention validity excludes padding; MEAP
422
+ # replacement rows are excluded from generator diagnostics. Token IDs alone
423
+ # are not a padding test (EOS may share the padding ID). JToK uses its existing
424
+ # validity mask. Let V be the remaining sampled rows and A(f)=mean_{r in V} f_r.
425
+ # Every row statistic below is reduced with A, except sample_valid_rows=|V|.
426
+ # Thus A(RMS(v_r)) is NOT RMS over all flattened activations. Empty V produces
427
+ # neutral zeros; consult sample_valid_rows before interpreting any metric.
428
+ # Deterministic position sampling is not an unbiased dataset estimate and can
429
+ # alias positional structure. The cache describes the last gradient-enabled
430
+ # microbatch/recomputation, not an optimizer-step, epoch or validation average.
431
+ # Arithmetic uses FP32 on detached, actually stored native tensors. Nonfinite
432
+ # values are not silently repaired; interpret RMS/angles with finite fractions.
433
+ #
434
+ # Leviathan metrics (prefix dynamics/leviathan/):
435
+ # Let e_r be the generated embedding, z_r the seed, and m_r the concatenation
436
+ # of stored CP product modes across heads/ranks. Define RMS(v)=sqrt(mean(v^2)),
437
+ # Z(v)=mean(1[v=0]), NF(v)=mean(1[not finite(v)]).
438
+ # embedding_rms = A(RMS(e)); tracks generator output scale.
439
+ # seed_rms = A(RMS(z)); tracks the shared seed representation's scale.
440
+ # product_mode_rms = A(RMS(m)); tracks the CP branch's numerical magnitude.
441
+ # embedding_nonfinite_fraction = A(NF(e)); flags invalid generated outputs.
442
+ # product_mode_zero_fraction = A(Z(m)); flags stored zero/underflow/collapse,
443
+ # without distinguishing their causes or the separate residual branch.
444
+ # product_mode_nonfinite_fraction = A(NF(m)); localizes invalid CP products.
445
+ # coordinate_saturation_fraction = A(mean_{head,d} 1[c<=.01 or c>=.99]),
446
+ # c=sigmoid(.5*(gamma*xhat+beta)). xhat/gamma/beta are native checkpoints;
447
+ # c is reconstructed only for sampled rows in FP32, before any activation
448
+ # cast. This is a boundary/conditioning proxy, not the exact stored
449
+ # coordinate's threshold test or proof that spline derivatives vanish.
450
+ # Seed, mode and embedding scales separate representation growth from product
451
+ # attenuation and projection effects. Their trends alone do not prove learning.
452
+ #
453
+ # JToK versus JToK-M: keep their distinct forward semantics when comparing.
454
+ # x is the incoming MLP update delta_m, not the full decoder hidden state;
455
+ # y is the update returned by this module, d=y-x, s is the learned scaler.
456
+ # With n(v)=v/(||v||_2+norm_eps), on valid rows:
457
+ # JToK: y=x*(1+s*n(S(z_tilde))) [elementwise modulation].
458
+ # JToK-M: y=x+a*s*n(sum_k p_k*S_{e_k}(z_tilde)) [additive residual],
459
+ # a=1/sqrt(2L), e_k=TopK(router_logits),
460
+ # p_k=sigmoid(logit_{e_k})/max(sum_j sigmoid(logit_{e_j}),norm_eps).
461
+ # Native stored output rounding is part of the observed d. The observer does
462
+ # not reconstruct an ideal gate/residual or change normalization/routing.
463
+ #
464
+ # Shared metrics (prefix dynamics/jtok|jtokm/layer_<index>/), eps=1e-12:
465
+ # input_rms=A(RMS(x)), output_rms=A(RMS(y)), update_rms=A(RMS(d));
466
+ # separate incoming scale, outgoing scale and effective intervention.
467
+ # relative_update_norm=A(||d||_2/max(||x||_2,eps)); unbounded effect size,
468
+ # sensitive to near-zero input. Read alongside zero_input_fraction.
469
+ # usage_index=A(||d||_2/max(||x||_2+||d||_2,eps)); bounded [0,1] for finite
470
+ # inputs, zero for an identity intervention. It measures visible effect,
471
+ # NOT probability of use, benefit, gradient flow or semantic relevance.
472
+ # input_output_cosine=A(clip(<x,y>/max(||x||_2*||y||_2,eps),-1,1));
473
+ # detects direction preservation/reversal separately from scale changes.
474
+ # update_alignment=A(clip(<x,d>/max(||x||_2*||d||_2,eps),-1,1));
475
+ # positive/negative values indicate reinforcement/cancellation of x.
476
+ # orthogonal_update_fraction=A(1-alignment_r^2 if ||d||_2>0 else 0);
477
+ # for nonzero, well-scaled x this is the fraction of d's energy
478
+ # orthogonal to x. With x=0,d!=0 the convention is 1 (no reference
479
+ # direction); epsilon-dominated angles are proxies, not exact geometry.
480
+ # It is the mean of per-row fractions, NOT 1-A(alignment)^2.
481
+ # changed_element_fraction=A(mean_h 1[y_h!=x_h]); visible dtype-level
482
+ # intervention, not an FP32-master update or usefulness test.
483
+ # zero_input_fraction=A(1[||x||_2=0]); qualifies norm ratios and angles.
484
+ # output_nonfinite_fraction=A(NF(y)); output numerical-health check.
485
+ # product_mode_zero_fraction=A(Z(m)), product_mode_nonfinite_fraction=A(NF(m));
486
+ # m contains all stored selected-expert CP modes (one expert for JToK).
487
+ # coordinate_saturation_fraction=A(mean_d 1[z_tilde_d<=.01 or >=.99]);
488
+ # observes the supplied coordinate, not the generator head coordinates.
489
+ # Jointly usage, alignment and orthogonal fraction distinguish a small/identity
490
+ # intervention, reinforcement/cancellation and direction-changing corrections.
491
+ # They cannot establish why a token needs that correction; causal usefulness
492
+ # requires a controlled ablation and task/loss evidence, not these proxies.
493
+ #
494
+ # Additional JToK-M routing metrics (not emitted for plain JToK):
495
+ # For selected stored weights p, b=sum_k p_k and pi_k=p_k/b when b>0 (else 0):
496
+ # selected_weight_entropy=A(-sum_k pi_k*log(pi_k)); entropy in nats over the
497
+ # selected K, not all E experts. K=1 is necessarily zero. Diagnostic pi
498
+ # renormalization does not modify the actual p used by the model.
499
+ # selected_max_weight=A(max_k p_k); concentration of actual mixture weights.
500
+ # selected_weight_sum=A(b), zero_selected_mass_fraction=A(1[b=0]); reveal
501
+ # epsilon-limited or underflowed mixture mass; entropy alone can hide it.
502
+ # Let f_e=sum_{r in V,k}1[e_k(r)=e]/(|V|*K), selection counts, NOT weight mass:
503
+ # expert_<e>_selection_fraction=f_e; sampled expert traffic.
504
+ # load_cv=std_population(f)/max(mean(f),eps); imbalance across experts.
505
+ # load_entropy_normalized=-sum_e f_e*log(max(f_e,1e-30))/max(log(E),eps);
506
+ # coverage/uniformity on [0,1] when V is nonempty and E>1; E=1 yields 0.
507
+ # active_experts=sum_e 1[f_e>0]; observed coverage, not causal contribution.
508
+ # These characterize routing collapse/specialization candidates, not expert
509
+ # quality. Low load entropy may be intentional specialization. Always compare
510
+ # with usage, update geometry, numerical health and train/validation loss.
511
+ #
512
+ # Cost/API contract: each observed forward samples at most sample_tokens rows,
513
+ # reduces on GPU and caches only compact detached FP32 scalars, never token IDs
514
+ # or activation vectors. There are extra GPU launches/workspaces, so cheap does
515
+ # not mean free. get_dynamics_metrics() exposes device scalars outside compiled
516
+ # forward; train.py batches their host transfer only when logging. Evaluation
517
+ # does not collect these training snapshots or relabel them as eval metrics.
518
+ try:
519
+ import triton
520
+ import triton.language as tl
521
+ except (ImportError, ModuleNotFoundError):
522
+ triton = tl = None
523
+
524
+ LEV_NAMES = (
525
+ "sample_valid_rows", "embedding_rms", "embedding_nonfinite_fraction",
526
+ "seed_rms", "coordinate_saturation_fraction", "product_mode_rms",
527
+ "product_mode_zero_fraction", "product_mode_nonfinite_fraction",
528
+ )
529
+ JTOK_NAMES = (
530
+ "sample_valid_rows", "input_rms", "output_rms", "update_rms",
531
+ "relative_update_norm", "input_output_cosine", "changed_element_fraction",
532
+ "product_mode_zero_fraction", "product_mode_nonfinite_fraction",
533
+ "coordinate_saturation_fraction", "selected_weight_entropy",
534
+ "selected_max_weight", "output_nonfinite_fraction", "zero_input_fraction",
535
+ "usage_index", "update_alignment", "orthogonal_update_fraction",
536
+ "selected_weight_sum", "zero_selected_mass_fraction",
537
+ )
538
+
539
+
540
+ if triton is not None:
541
+ @triton.jit
542
+ def _dynamics_reduce_kernel(rows, output, S: tl.constexpr, C: tl.constexpr,
543
+ BS: tl.constexpr):
544
+ col = tl.program_id(0)
545
+ r = tl.arange(0, BS)
546
+ valid = tl.load(rows + r * C, r < S, 0)
547
+ count = tl.sum(valid, 0)
548
+ values = tl.load(rows + r * C + col, r < S, 0)
549
+ value = tl.sum(tl.where(valid > 0, values, 0), 0) / tl.maximum(count, 1)
550
+ tl.store(output + col, tl.where(col == 0, count, value))
551
+
552
+ @triton.jit
553
+ def _jtok_dynamics_rows_kernel(delta, out, z, modes, experts, weights, valid,
554
+ rows, N: tl.constexpr, S: tl.constexpr,
555
+ H: tl.constexpr, D: tl.constexpr, M: tl.constexpr,
556
+ K: tl.constexpr, E: tl.constexpr, BH: tl.constexpr,
557
+ BD: tl.constexpr, BM: tl.constexpr,
558
+ BK: tl.constexpr, HAS_MASK: tl.constexpr):
559
+ sample = tl.program_id(0)
560
+ row = sample * N // S
561
+ is_valid = tl.full((), True, tl.int1)
562
+ if HAS_MASK:
563
+ is_valid = tl.load(valid + row)
564
+ h = tl.arange(0, BH)
565
+ x = tl.load(delta + row * H + h, h < H, 0).to(tl.float32)
566
+ y = tl.load(out + row * H + h, h < H, 0).to(tl.float32)
567
+ diff = y - x
568
+ xn = tl.sqrt(tl.sum(x * x, 0))
569
+ yn = tl.sqrt(tl.sum(y * y, 0))
570
+ dn = tl.sqrt(tl.sum(diff * diff, 0))
571
+ cosine = tl.minimum(1., tl.maximum(-1., tl.sum(x * y, 0) / tl.maximum(xn * yn, 1e-12)))
572
+ d = tl.arange(0, BD)
573
+ coord = tl.load(z + row * D + d, d < D, 0).to(tl.float32)
574
+ saturation = tl.sum(((coord <= .01) | (coord >= .99)) & (d < D), 0) / D
575
+ m = tl.arange(0, BM)
576
+ mode = tl.load(modes + row * K * M + m, m < K * M, 0).to(tl.float32)
577
+ mz = tl.sum((mode == 0) & (m < K * M), 0) / (K * M)
578
+ mn = tl.sum(((mode != mode) | (tl.abs(mode) == float('inf'))) & (m < K * M), 0) / (K * M)
579
+ k = tl.arange(0, BK)
580
+ p = tl.load(weights + row * K + k, k < K, 0).to(tl.float32)
581
+ mass = tl.sum(p, 0)
582
+ probability = p / tl.where(mass > 0, mass, 1.)
583
+ entropy = -tl.sum(tl.where(probability > 0,
584
+ probability * tl.log(tl.maximum(probability, 1e-30)), 0), 0)
585
+ max_weight = tl.max(p, 0)
586
+ nonfinite = tl.sum(((y != y) | (tl.abs(y) == float('inf'))) & (h < H), 0) / H
587
+ changed = tl.sum((y != x) & (h < H), 0) / H
588
+ alignment = tl.minimum(1., tl.maximum(-1., tl.sum(x * diff, 0) / tl.maximum(xn * dn, 1e-12)))
589
+ orthogonal = tl.where(dn > 0, tl.maximum(0., 1. - alignment * alignment), 0.)
590
+ C: tl.constexpr = 19 + E
591
+ base = rows + sample * C
592
+ tl.store(base, is_valid.to(tl.float32))
593
+ tl.store(base + 1, xn / tl.sqrt(float(H)))
594
+ tl.store(base + 2, yn / tl.sqrt(float(H)))
595
+ tl.store(base + 3, dn / tl.sqrt(float(H)))
596
+ tl.store(base + 4, dn / tl.maximum(xn, 1e-12))
597
+ tl.store(base + 5, cosine)
598
+ tl.store(base + 6, changed)
599
+ tl.store(base + 7, mz)
600
+ tl.store(base + 8, mn)
601
+ tl.store(base + 9, saturation)
602
+ tl.store(base + 10, entropy)
603
+ tl.store(base + 11, max_weight)
604
+ tl.store(base + 12, nonfinite)
605
+ tl.store(base + 13, (xn == 0).to(tl.float32))
606
+ tl.store(base + 14, dn / tl.maximum(xn + dn, 1e-12))
607
+ tl.store(base + 15, alignment)
608
+ tl.store(base + 16, orthogonal)
609
+ tl.store(base + 17, mass)
610
+ tl.store(base + 18, (mass == 0).to(tl.float32))
611
+ route = tl.load(experts + row * K + k, k < K, -1)
612
+ for expert in range(E):
613
+ fraction = tl.sum((route == expert) & (k < K), 0) / K
614
+ tl.store(base + 19 + expert, fraction)
615
+
616
+ @triton.jit
617
+ def _lev_dynamics_rows_kernel(embedding, seed, xhat, norm_weight, norm_bias,
618
+ modes, valid, rows, N: tl.constexpr, S: tl.constexpr,
619
+ H: tl.constexpr, D: tl.constexpr, HEADS: tl.constexpr,
620
+ R: tl.constexpr, BH: tl.constexpr, BD: tl.constexpr,
621
+ BM: tl.constexpr, HAS_MASK: tl.constexpr):
622
+ sample = tl.program_id(0)
623
+ row = sample * N // S
624
+ is_valid = tl.full((), True, tl.int1)
625
+ if HAS_MASK:
626
+ is_valid = tl.load(valid + row)
627
+ h = tl.arange(0, BH)
628
+ value = tl.load(embedding + row * H + h, h < H, 0).to(tl.float32)
629
+ erms = tl.sqrt(tl.sum(value * value, 0) / H)
630
+ en = tl.sum(((value != value) | (tl.abs(value) == float('inf'))) & (h < H), 0) / H
631
+ d = tl.arange(0, BD)
632
+ z = tl.load(seed + row * D + d, d < D, 0).to(tl.float32)
633
+ zrms = tl.sqrt(tl.sum(z * z, 0) / D)
634
+ saturated = tl.full((), 0., tl.float32)
635
+ for head in range(HEADS):
636
+ x = tl.load(xhat + (head * N + row) * D + d, d < D, 0).to(tl.float32)
637
+ w = tl.load(norm_weight + head * D + d, d < D, 0).to(tl.float32)
638
+ b = tl.load(norm_bias + head * D + d, d < D, 0).to(tl.float32)
639
+ coordinate = tl.sigmoid(.5 * (x * w + b))
640
+ saturated += tl.sum(((coordinate <= .01) | (coordinate >= .99)) & (d < D), 0)
641
+ m = tl.arange(0, BM)
642
+ mode = tl.load(modes + row * HEADS * R + m, m < HEADS * R, 0).to(tl.float32)
643
+ mrms = tl.sqrt(tl.sum(mode * mode, 0) / (HEADS * R))
644
+ mz = tl.sum((mode == 0) & (m < HEADS * R), 0) / (HEADS * R)
645
+ mn = tl.sum(((mode != mode) | (tl.abs(mode) == float('inf'))) & (m < HEADS * R), 0) / (HEADS * R)
646
+ base = rows + sample * 8
647
+ tl.store(base, is_valid.to(tl.float32))
648
+ tl.store(base + 1, erms)
649
+ tl.store(base + 2, en)
650
+ tl.store(base + 3, zrms)
651
+ tl.store(base + 4, saturated / (HEADS * D))
652
+ tl.store(base + 5, mrms)
653
+ tl.store(base + 6, mz)
654
+ tl.store(base + 7, mn)
655
+
656
+
657
+ def _check_sampling(max_samples: int) -> None:
658
+ validate_dynamics_samples(max_samples)
659
+ if max_samples == 0:
660
+ raise ValueError("diagnostics max_samples must be in [1, 4096]")
661
+
662
+
663
+ def validate_dynamics_samples(samples: int) -> None:
664
+ """Zero disables collection; reject invalid budgets before native work."""
665
+ if type(samples) is not int or not 0 <= samples <= 4096:
666
+ raise ValueError("dynamics_samples must be an integer in [0, 4096]")
667
+
668
+
669
+ def _check_tensors(primary: torch.Tensor, *others: torch.Tensor) -> None:
670
+ for value in (primary, *others):
671
+ if value.device != primary.device or not value.is_contiguous():
672
+ raise ValueError("diagnostics tensors must be contiguous on the same device")
673
+ if value.dtype not in (torch.float16, torch.bfloat16, torch.float32):
674
+ raise TypeError("diagnostics floating tensors require FP16, BF16 or FP32")
675
+
676
+
677
+ def _check_mask(valid: torch.Tensor, n: int, device: torch.device) -> None:
678
+ if (valid.dtype != torch.bool or valid.device != device or valid.ndim != 1
679
+ or valid.numel() not in (0, n) or not valid.is_contiguous()):
680
+ raise ValueError("diagnostics mask must be a contiguous bool vector of length zero or N on the input device")
681
+
682
+
683
+ def _reduce(rows: torch.Tensor, columns: int) -> torch.Tensor:
684
+ output = torch.empty(columns, device=rows.device, dtype=torch.float32)
685
+ _dynamics_reduce_kernel[(columns,)](
686
+ rows, output, rows.shape[0], columns,
687
+ triton.next_power_of_2(rows.shape[0]), num_warps=4)
688
+ return output
689
+
690
+
691
+ @torch.library.custom_op("neollm_dynamics::jtok_dynamics", mutates_args=())
692
+ def jtok_dynamics(delta: torch.Tensor, output: torch.Tensor, z: torch.Tensor,
693
+ modes: torch.Tensor, experts: torch.Tensor, weights: torch.Tensor,
694
+ valid: torch.Tensor, num_experts: int, max_samples: int) -> torch.Tensor:
695
+ _check_sampling(max_samples)
696
+ if type(num_experts) is not int or not 1 <= num_experts <= 1024:
697
+ raise ValueError("diagnostics num_experts must be an integer in [1, 1024]")
698
+ if delta.ndim != 2 or delta.shape[1] == 0 or output.shape != delta.shape:
699
+ raise ValueError("diagnostics input/output must have shape [N, H] with H > 0")
700
+ n, hidden = delta.shape
701
+ if z.ndim != 2 or z.shape[0] != n or z.shape[1] == 0:
702
+ raise ValueError("diagnostics coordinates must have shape [N, D] with D > 0")
703
+ if (modes.ndim != 3 or modes.shape[0] != n or modes.shape[2] == 0
704
+ or not 1 <= modes.shape[1] <= num_experts):
705
+ raise ValueError("diagnostics modes must have shape [N, K, M] with 1 <= K <= E and M > 0")
706
+ if experts.shape != modes.shape[:2] or weights.shape != experts.shape:
707
+ raise ValueError("diagnostics routes/weights must have shape [N, K]")
708
+ if (experts.dtype not in (torch.int32, torch.int64)
709
+ or experts.device != delta.device or not experts.is_contiguous()):
710
+ raise ValueError("diagnostics routes must be contiguous integer tensors on the input device")
711
+ _check_tensors(delta, output, z, modes, weights)
712
+ _check_mask(valid, n, delta.device)
713
+ if triton is None or not delta.is_cuda:
714
+ raise RuntimeError("native JToK diagnostics require CUDA/Triton")
715
+ samples = min(n, max_samples)
716
+ columns = len(JTOK_NAMES) + num_experts
717
+ if not samples:
718
+ return torch.zeros(columns, device=delta.device, dtype=torch.float32)
719
+ rows = torch.empty((samples, columns), device=delta.device, dtype=torch.float32)
720
+ _jtok_dynamics_rows_kernel[(samples,)](
721
+ delta, output, z, modes, experts, weights, valid, rows, n, samples,
722
+ hidden, z.shape[1], modes.shape[2], modes.shape[1], num_experts,
723
+ triton.next_power_of_2(hidden), triton.next_power_of_2(z.shape[1]),
724
+ triton.next_power_of_2(modes.shape[1] * modes.shape[2]),
725
+ triton.next_power_of_2(modes.shape[1]), valid.numel() != 0,
726
+ num_warps=4)
727
+ return _reduce(rows, columns)
728
+
729
+
730
+ @jtok_dynamics.register_fake
731
+ def _jtok_dynamics_fake(delta, output, z, modes, experts, weights, valid,
732
+ num_experts, max_samples):
733
+ return torch.empty(len(JTOK_NAMES) + num_experts, device=delta.device, dtype=torch.float32)
734
+
735
+
736
+ @torch.library.custom_op("neollm_dynamics::leviathan_dynamics", mutates_args=())
737
+ def leviathan_dynamics(embedding: torch.Tensor, seed: torch.Tensor, xhat: torch.Tensor,
738
+ norm_weight: torch.Tensor, norm_bias: torch.Tensor,
739
+ modes: torch.Tensor, valid: torch.Tensor, max_samples: int) -> torch.Tensor:
740
+ _check_sampling(max_samples)
741
+ if embedding.ndim < 2 or embedding.shape[-1] == 0:
742
+ raise ValueError("diagnostics embedding must have shape [..., H] with H > 0")
743
+ n, hidden = embedding.numel() // embedding.shape[-1], embedding.shape[-1]
744
+ if (xhat.ndim != 3 or xhat.shape[0] == 0 or xhat.shape[1] != n
745
+ or xhat.shape[2] == 0):
746
+ raise ValueError("diagnostics xhat must have shape [heads, N, D] with positive heads and D")
747
+ heads, _, d = xhat.shape
748
+ if seed.shape != (n, d) or norm_weight.shape != (heads, d) or norm_bias.shape != (heads, d):
749
+ raise ValueError("diagnostics seed/norm metadata do not match xhat")
750
+ if modes.ndim != 3 or modes.shape[:2] != (n, heads) or modes.shape[2] == 0:
751
+ raise ValueError("diagnostics modes must have shape [N, heads, rank] with rank > 0")
752
+ _check_tensors(embedding, seed, xhat, norm_weight, norm_bias, modes)
753
+ _check_mask(valid, n, embedding.device)
754
+ if triton is None or not embedding.is_cuda:
755
+ raise RuntimeError("native Leviathan diagnostics require CUDA/Triton")
756
+ samples = min(n, max_samples)
757
+ if not samples:
758
+ return torch.zeros(len(LEV_NAMES), device=embedding.device, dtype=torch.float32)
759
+ rows = torch.empty((samples, len(LEV_NAMES)), device=embedding.device, dtype=torch.float32)
760
+ rank = modes.shape[-1]
761
+ _lev_dynamics_rows_kernel[(samples,)](
762
+ embedding, seed, xhat, norm_weight, norm_bias, modes, valid, rows,
763
+ n, samples, hidden, d, heads, rank, triton.next_power_of_2(hidden),
764
+ triton.next_power_of_2(d), triton.next_power_of_2(heads * rank),
765
+ valid.numel() != 0, num_warps=4)
766
+ return _reduce(rows, len(LEV_NAMES))
767
+
768
+
769
+ @leviathan_dynamics.register_fake
770
+ def _leviathan_dynamics_fake(embedding, seed, xhat, norm_weight, norm_bias,
771
+ modes, valid, max_samples):
772
+ return torch.empty(len(LEV_NAMES), device=embedding.device, dtype=torch.float32)
773
+
774
+
775
+ def unpack_jtok_dynamics(packed: torch.Tensor, num_experts: int, *, mixture: bool):
776
+ """Expose device scalars; the caller chooses when to transfer to its logger."""
777
+ result = dict(zip(JTOK_NAMES, packed[:len(JTOK_NAMES)].unbind()))
778
+ if not mixture:
779
+ result.pop("selected_weight_entropy")
780
+ result.pop("selected_max_weight")
781
+ result.pop("selected_weight_sum")
782
+ result.pop("zero_selected_mass_fraction")
783
+ return result
784
+ fractions = packed[len(JTOK_NAMES):]
785
+ mean = fractions.mean()
786
+ result["load_cv"] = fractions.std(unbiased=False) / mean.clamp_min(1e-12)
787
+ result["load_entropy_normalized"] = (
788
+ -(fractions * fractions.clamp_min(1e-30).log()).sum()
789
+ / torch.tensor(float(num_experts), device=packed.device).log().clamp_min(1e-12))
790
+ result["active_experts"] = (fractions > 0).sum().float()
791
+ for expert, value in enumerate(fractions.unbind()):
792
+ result[f"expert_{expert}_selection_fraction"] = value
793
+ return result
794
+
795
+
796
+
797
+
798
  class LeviathanGenerator(nn.Module):
799
  """
800
  Learned Embedding Vectorization (LEV) layer - Leviathan input embedding
 
843
  # is unavailable; callers can also disable the optional route per
844
  # instance without changing the model/state-dict contract.
845
  self.use_leviathan_triton = bool(_LEV_KERNEL_AVAILABLE)
846
+ self.dynamics_samples = (
847
+ int(getattr(config, "dynamics_sample_tokens", 64))
848
+ if bool(getattr(config, "dynamics_metrics_enabled", False)) else 0
849
+ )
850
+ self._last_dynamics_metrics = None
851
  vocab_size = config.vocab_size
852
  hidden_size = config.hidden_size
853
  self.d_seed = config.generator_d_seed
 
1036
  prod_sign = 1.0 - 2.0 * (num_neg % 2).float() # [N, krank]
1037
  return prod_sign * torch.exp(log_mag) # [N, krank]
1038
 
1039
+ def _record_dynamics(self, token_ids, embedding, checkpoints, valid_mask, mask_token_id):
1040
+ seed, xhat, modes = checkpoints
1041
+ diag_mask = (token_ids.new_empty((0,), dtype=torch.bool) if valid_mask is None
1042
+ else valid_mask.reshape(-1).contiguous())
1043
+ if mask_token_id is not None:
1044
+ unmasked = token_ids.reshape(-1).ne(int(mask_token_id))
1045
+ diag_mask = unmasked if diag_mask.numel() == 0 else diag_mask & unmasked
1046
+ packed = leviathan_dynamics(
1047
+ embedding.detach(), seed, xhat, self.head_norm_weight.detach(),
1048
+ self.head_norm_bias.detach(), modes, diag_mask, self.dynamics_samples)
1049
+ self._last_dynamics_metrics = packed.detach().clone()
1050
+
1051
  # Keep this method traceable. The compiler-safe LEV entry point registers
1052
  # an opaque custom op tagged cudagraph_unsafe: Dynamo keeps the LEV call
1053
  # in the compiled graph while Inductor excludes only the unsafe Triton
 
1060
  meap_mask_embedding: Optional[torch.Tensor] = None,
1061
  meap_mask_token_id: Optional[int] = None,
1062
  return_jtok_geometry: bool = False,
1063
+ dynamics_valid_mask: Optional[torch.Tensor] = None,
1064
  ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
1065
  """
1066
  Generate embeddings from discrete token indices.
 
1075
  where ``z_tilde`` is the shared differentiable Leviathan/JTok
1076
  coordinate with shape ``[*token_ids.shape, d_seed]``.
1077
  """
1078
+ collect_dynamics = bool(self.dynamics_samples and self.training and torch.is_grad_enabled())
1079
+ dynamics_kwargs = (
1080
+ {"return_dynamics_checkpoints": True}
1081
+ if collect_dynamics else {}
1082
+ )
1083
  has_meap = (
1084
  meap_mask_embedding is not None or meap_mask_token_id is not None
1085
  )
 
1119
  )
1120
  require_leviathan_triton = bool(
1121
  getattr(self.config, "_require_leviathan_triton", False)
1122
+ ) or collect_dynamics
1123
  if require_leviathan_triton and not leviathan_kernel_ready:
1124
  raise RuntimeError(
1125
  "This run requires the Triton Leviathan kernel, but the "
 
1164
  "but the installed CCE package lacks the geometry "
1165
  "adapter."
1166
  )
1167
+ result = _leviathan_embedding_with_seed_compiler_safe(
1168
  token_ids,
1169
  params,
1170
  self.config,
1171
  self.knot_grid,
1172
  mask_embedding=meap_mask_embedding,
1173
  mask_token_id=meap_mask_token_id,
1174
+ **dynamics_kwargs,
1175
  )
1176
+ if collect_dynamics:
1177
+ embedding, seed, checkpoints = result
1178
+ self._record_dynamics(token_ids, embedding, checkpoints,
1179
+ dynamics_valid_mask, meap_mask_token_id)
1180
+ else:
1181
+ embedding, seed = result
1182
  embedding = embedding.reshape(*token_ids.shape, self.hidden_size)
1183
  return embedding, self._jtok_coordinate_from_seed(seed)
1184
+ result = _leviathan_embedding_compiler_safe(
1185
  token_ids,
1186
  params,
1187
  self.config,
1188
  self.knot_grid,
1189
  mask_embedding=meap_mask_embedding,
1190
  mask_token_id=meap_mask_token_id,
1191
+ **dynamics_kwargs,
1192
  )
1193
+ if collect_dynamics:
1194
+ embedding, checkpoints = result
1195
+ self._record_dynamics(token_ids, embedding, checkpoints,
1196
+ dynamics_valid_mask, meap_mask_token_id)
1197
+ return embedding
1198
+ return result
1199
  except (TypeError, ValueError, AttributeError):
1200
  # Only validation/configuration failures use the reference
1201
  # implementation. A CUDA launch error must propagate: doing
 
1299
  super().__init__()
1300
  self.layer_idx = int(layer_idx)
1301
  self.use_mixture = bool(getattr(config, "use_jtokm", False))
1302
+ self.dynamics_samples = (
1303
+ int(getattr(config, "dynamics_sample_tokens", 64))
1304
+ if bool(getattr(config, "dynamics_metrics_enabled", False)) else 0
1305
+ )
1306
+ self._last_dynamics_metrics = None
1307
  self.jtok_kernel_backend = str(
1308
  getattr(config, "jtok_kernel_backend", "torch")
1309
  ).strip().lower()
 
1311
  raise ValueError(
1312
  "`jtok_kernel_backend` must be either `torch` or `triton`."
1313
  )
1314
+ if self.dynamics_samples and self.jtok_kernel_backend != "triton":
1315
+ raise RuntimeError("Native dynamics metrics require the strict Triton JTok backend")
1316
  if self.jtok_kernel_backend == "triton" and not _JTOK_KERNEL_AVAILABLE:
1317
  raise RuntimeError(
1318
  "`jtok_kernel_backend='triton'` was requested, but the "
 
1540
  Optional[Tuple[torch.Tensor, torch.Tensor, torch.Tensor]],
1541
  ]:
1542
  """Apply JTok/JTok-M and optionally return training-only router stats."""
1543
+ collect_dynamics = bool(self.dynamics_samples and self.training and torch.is_grad_enabled())
1544
  if self.jtok_kernel_backend == "triton":
1545
  # Dispatch only: this branch does not change compiler mode,
1546
  # CUDA-graph policy, training state, optimizer behavior, or any
 
1553
  "JTok Triton adapter became unavailable after module "
1554
  "construction; refusing an implicit Torch fallback."
1555
  )
1556
+ result = _apply_neollm_jtok(
1557
  self,
1558
  delta_m,
1559
  z_tilde,
 
1561
  valid_mask=valid_mask,
1562
  compute_aux=bool(compute_aux and self.use_mixture),
1563
  backend="triton",
1564
+ **({"return_dynamics_checkpoints": True} if collect_dynamics else {}),
1565
+ )
1566
+ if collect_dynamics:
1567
+ output, aux, checkpoints = result
1568
+ modes, routes, weights = checkpoints
1569
+ n = delta_m.numel() // self.hidden_size
1570
+ mask = (delta_m.new_empty((0,), dtype=torch.bool) if valid_mask is None
1571
+ else valid_mask.reshape(-1).to(dtype=torch.bool).contiguous())
1572
+ packed = jtok_dynamics(
1573
+ delta_m.detach().reshape(n, self.hidden_size).contiguous(),
1574
+ output.detach().reshape(n, self.hidden_size).contiguous(),
1575
+ z_tilde.detach().reshape(n, self.d_seed).contiguous(), modes, routes, weights,
1576
+ mask, self.num_experts, self.dynamics_samples)
1577
+ self._last_dynamics_metrics = packed.detach().clone()
1578
+ return output, aux
1579
+ return result
1580
 
1581
  orig_shape = delta_m.shape
1582
  N = delta_m.numel() // self.hidden_size
 
1738
  self.down_proj = nn.Linear(config.hidden_size, config.stack_d_model)
1739
  self.up_proj = nn.Linear(config.stack_d_model, config.hidden_size)
1740
  self.action_head = nn.Linear(config.stack_d_model, 3 * self.num_mem_heads)
1741
+ # A scalar bias is exactly softmax-invariant because the same value is
1742
+ # added to every stack slot. Keeping it would create a parameter tensor
1743
+ # with an analytically zero gradient, which is especially undesirable for
1744
+ # per-tensor GradientStabilizer scaling.
1745
+ self.gate_proj = nn.Linear(self.stack_dim, 1, bias=False)
1746
  self.res_weight = nn.Parameter(torch.ones(1))
1747
 
1748
  self.cache_size = getattr(config, "stack_memory_cache_size", 2048)
 
1767
  [mask[:, :, :, 1:], torch.zeros_like(mask[:, :, :, :1])], dim=3
1768
  )
1769
 
1770
+ # Keep stack values in the model dtype (normally BF16), but keep the
1771
+ # differentiable control mask in FP32. The latter is important because
1772
+ # the global-read mask bias can be extremely large (e.g. -1e9).
1773
+ action_weights_stack = actions.to(stack.dtype).unsqueeze(-1).unsqueeze(-1)
1774
+ action_weights_mask = actions.to(mask.dtype).unsqueeze(-1)
1775
  stacks = torch.stack([push_stack, pop_stack, stack], dim=3)
1776
  masks = torch.stack([push_mask, pop_mask, mask], dim=3)
1777
 
1778
+ new_stack = (stacks * action_weights_stack).sum(dim=3)
1779
+ new_mask = (masks * action_weights_mask).sum(dim=3)
1780
  return new_stack, new_mask
1781
 
1782
+ def forward(self, hidden_states, stack, mask, token_mask=None):
1783
  batch_size, seq_len, _ = hidden_states.shape
1784
 
1785
+ # The stack occupancy is a differentiable control state. Keep it in
1786
+ # FP32 across layers/callers even when model activations are BF16.
1787
+ mask = mask.float()
1788
+
1789
  # OLMo stack_v3 value/action path: project the complete hidden state into
1790
  # the total low-rank stack width, then split that representation by head.
1791
  new_hidden_states = self.down_proj(hidden_states)
1792
 
1793
  action_logits = self.action_head(new_hidden_states) / math.sqrt(self.stack_dim)
1794
+ # Stack actions are control probabilities. Evaluate softmax in FP32 so
1795
+ # the recurrent mask is never quantized to BF16 before later layers use
1796
+ # it. Gradients still flow normally into the BF16 action_head weights.
1797
  actions = F.softmax(
1798
+ action_logits.float().reshape(
1799
+ batch_size, seq_len, self.num_mem_heads, 3
1800
+ ),
1801
  dim=-1,
1802
  )
1803
 
 
1807
 
1808
  new_stack, new_mask = self._vectorized_update(stack, mask, actions, k_values)
1809
 
1810
+ # The global-read control path is evaluated explicitly in FP32. This
1811
+ # preserves the original differentiable graph (no detach) while avoiding
1812
+ # BF16 quantization of both the soft mask and the very large mask bias.
1813
+ stack_fp32 = new_stack.float()
1814
+ gate_scores = F.linear(
1815
+ stack_fp32,
1816
+ self.gate_proj.weight.float(),
1817
+ bias=None,
1818
+ ).squeeze(-1)
1819
+ gate_weights = F.softmax(
1820
+ gate_scores + (1.0 - new_mask) * -320,
1821
+ dim=-1,
1822
+ )
1823
 
1824
+ memory_output = (stack_fp32 * gate_weights.unsqueeze(-1)).sum(dim=3)
1825
+ memory_output = memory_output.to(new_hidden_states.dtype)
1826
  memory_output = memory_output.reshape(batch_size, seq_len, -1)
1827
  memory_output = self.up_proj(memory_output)
1828
 
1829
  output = memory_output * self.res_weight + hidden_states
1830
 
1831
+ # Compact StackTrans statistics. Entropy is evaluated in FP32 for
1832
+ # numerical stability; only its scalar mean keeps autograd because it is
1833
+ # used by the auxiliary objective. All other diagnostics are detached
1834
+ # before they leave this layer so monitoring does not retain extra graphs.
1835
+ actions_fp32 = actions.float()
1836
+ action_entropy_map = -(
1837
+ actions_fp32 * torch.log(actions_fp32.clamp_min(1e-12))
1838
+ ).sum(dim=-1) # [B,S,H]
1839
+ action_max_map = actions_fp32.max(dim=-1).values
1840
+
1841
+ gate_fp32 = gate_weights.float()
1842
+ gate_entropy_map = -(
1843
+ gate_fp32 * torch.log(gate_fp32.clamp_min(1e-12))
1844
+ ).sum(dim=-1) # [B,S,H]
1845
+ gate_max_map = gate_fp32.max(dim=-1).values
1846
+ gate_effective_slots_map = gate_entropy_map.exp()
1847
+
1848
+ mask_fp32 = new_mask.float()
1849
+ expected_depth_map = mask_fp32.sum(dim=-1) # [B,S,H]
1850
+ mask_fill_map = expected_depth_map / float(self.stack_slots)
1851
+ top_occupancy_map = mask_fp32[..., 0]
1852
+
1853
+ if token_mask is None:
1854
+ action_entropy_mean = action_entropy_map.mean()
1855
+ action_probs_mean = actions_fp32.mean(dim=(0, 1, 2))
1856
+ action_max_mean = action_max_map.mean()
1857
+ expected_depth_mean = expected_depth_map.mean()
1858
+ mask_fill_mean = mask_fill_map.mean()
1859
+ top_occupancy_mean = top_occupancy_map.mean()
1860
+ gate_entropy_mean = gate_entropy_map.mean()
1861
+ gate_max_mean = gate_max_map.mean()
1862
+ gate_effective_slots_mean = gate_effective_slots_map.mean()
1863
+ else:
1864
+ valid = token_mask.to(
1865
+ device=hidden_states.device, dtype=torch.float32
1866
+ ).unsqueeze(-1) # [B,S,1]
1867
+ denom = (valid.sum() * float(self.num_mem_heads)).clamp_min(1.0)
1868
+ action_entropy_mean = (action_entropy_map * valid).sum() / denom
1869
+ action_probs_mean = (
1870
+ actions_fp32 * valid.unsqueeze(-1)
1871
+ ).sum(dim=(0, 1, 2)) / denom
1872
+ action_max_mean = (action_max_map * valid).sum() / denom
1873
+ expected_depth_mean = (expected_depth_map * valid).sum() / denom
1874
+ mask_fill_mean = (mask_fill_map * valid).sum() / denom
1875
+ top_occupancy_mean = (top_occupancy_map * valid).sum() / denom
1876
+ gate_entropy_mean = (gate_entropy_map * valid).sum() / denom
1877
+ gate_max_mean = (gate_max_map * valid).sum() / denom
1878
+ gate_effective_slots_mean = (
1879
+ gate_effective_slots_map * valid
1880
+ ).sum() / denom
1881
+
1882
+ stack_metrics = torch.stack(
1883
+ (
1884
+ action_entropy_mean,
1885
+ action_probs_mean[0].detach(),
1886
+ action_probs_mean[1].detach(),
1887
+ action_probs_mean[2].detach(),
1888
+ action_max_mean.detach(),
1889
+ expected_depth_mean.detach(),
1890
+ mask_fill_mean.detach(),
1891
+ top_occupancy_mean.detach(),
1892
+ gate_entropy_mean.detach(),
1893
+ gate_max_mean.detach(),
1894
+ gate_effective_slots_mean.detach(),
1895
+ )
1896
+ )
1897
+
1898
  if self.training and self.enable_cache:
1899
  self._update_cache(k_values.detach(), actions.detach())
1900
 
1901
  # Preserve the complete token-local stack for the next decoder layer.
1902
+ return output, new_stack, new_mask, stack_metrics
1903
 
1904
  def _update_cache(self, k_values, actions):
1905
  seq_len = k_values.shape[1]
 
1923
  if mask.ndim == 3:
1924
  mask = mask.unsqueeze(1)
1925
 
1926
+ output, new_stack, new_mask, _stack_metrics = self.forward(
1927
  hidden_state.unsqueeze(1), stack, mask
1928
  )
1929
  return output.squeeze(1), new_stack.squeeze(1), new_mask.squeeze(1)
 
6034
  hidden_states: torch.Tensor,
6035
  stack_memory: Optional[torch.Tensor],
6036
  stack_memory_mask: Optional[torch.Tensor],
6037
+ token_mask: Optional[torch.Tensor] = None,
6038
+ ) -> Tuple[
6039
+ torch.Tensor,
6040
+ Optional[torch.Tensor],
6041
+ Optional[torch.Tensor],
6042
+ Optional[torch.Tensor],
6043
+ ]:
6044
  """
6045
  Apply this layer's differentiable StackMemory module before attention.
6046
 
 
6050
  attention implementation remains unchanged.
6051
  """
6052
  if (not self.use_stack_memory) or self.stack_memory is None:
6053
+ return hidden_states, stack_memory, stack_memory_mask, None
6054
  if stack_memory is None or stack_memory_mask is None:
6055
  raise ValueError(
6056
  "StackMemory is enabled, but stack_memory/stack_memory_mask "
6057
  "were not initialized by NeoLLMModel.forward."
6058
  )
6059
+ return self.stack_memory(
6060
+ hidden_states, stack_memory, stack_memory_mask, token_mask=token_mask
6061
+ )
6062
 
6063
  def _attn_res(
6064
  self,
 
6179
  position_embeddings: tuple[torch.Tensor, torch.Tensor],
6180
  stack_memory: Optional[torch.Tensor] = None,
6181
  stack_memory_mask: Optional[torch.Tensor] = None,
6182
+ stack_token_mask: Optional[torch.Tensor] = None,
6183
  attention_mask: Optional[torch.Tensor] = None,
6184
  first_layer_fan: Optional[torch.Tensor] = None,
6185
  output_attentions: Optional[bool] = False,
 
6203
  if self.siamese_attn_input_norm is not None:
6204
  h_attn = self.siamese_attn_input_norm(h_attn)
6205
 
6206
+ stack_metrics = None
6207
  if self.use_stack_memory:
6208
+ h_attn, stack_memory, stack_memory_mask, stack_metrics = (
6209
+ self.apply_stack_memory(
6210
+ h_attn,
6211
+ stack_memory,
6212
+ stack_memory_mask,
6213
+ token_mask=stack_token_mask,
6214
+ )
6215
  )
6216
  jtok_router_state = h_attn
6217
 
 
6268
 
6269
  outputs = (x_next, y_next)
6270
  if self.use_stack_memory:
6271
+ outputs += (stack_memory, stack_memory_mask, stack_metrics)
6272
  if jtok_aux_stats is not None:
6273
  outputs += jtok_aux_stats
6274
  if output_attentions:
 
7199
  )
7200
  jtok_z_tilde = None
7201
 
7202
+ # Diagnostic validity is separate from attention/loss/JTok masks.
7203
+ # In particular, EOS may also be the padding ID; never infer padding
7204
+ # from IDs. Exclude MEAP replacements even if SBE owns the epilogue.
7205
+ dynamics_kwargs = {}
7206
+ if (bool(getattr(self.config, "dynamics_metrics_enabled", False))
7207
+ and self.training and torch.is_grad_enabled() and input_ids is not None):
7208
+ if attention_mask is not None:
7209
+ if attention_mask.ndim != 2 or attention_mask.shape != input_ids.shape:
7210
+ raise ValueError("Dynamics collection requires an explicit 2D attention validity mask")
7211
+ dynamics_mask = attention_mask.to(dtype=torch.bool)
7212
+ else:
7213
+ dynamics_mask = torch.ones_like(input_ids, dtype=torch.bool)
7214
+ if meap_active:
7215
+ dynamics_mask = dynamics_mask & input_ids.ne(int(self.config.meap_mask_token_id))
7216
+ dynamics_kwargs = {"dynamics_valid_mask": dynamics_mask}
7217
+
7218
  # ── Embedding stage ────────────────────────────────────────────────
7219
  meap_mask_embedding = None
7220
  meap_mask_token_id = None
 
7243
  meap_mask_embedding=meap_mask_embedding,
7244
  meap_mask_token_id=meap_mask_token_id,
7245
  return_jtok_geometry=True,
7246
+ **dynamics_kwargs,
7247
  )
7248
  else:
7249
  # Keep the pre-JTok call signature for custom or
 
7252
  input_ids,
7253
  meap_mask_embedding=meap_mask_embedding,
7254
  meap_mask_token_id=meap_mask_token_id,
7255
+ **dynamics_kwargs,
7256
  )
7257
  else:
7258
  if use_jtok:
7259
  generator_result = self.token_generator(
7260
  input_ids,
7261
  return_jtok_geometry=True,
7262
+ **dynamics_kwargs,
7263
  )
7264
  else:
7265
+ generator_result = self.token_generator(input_ids, **dynamics_kwargs)
7266
  if use_jtok:
7267
  inputs_embeds, jtok_z_tilde = generator_result
7268
  else:
 
7401
  .to(dtype=hidden_states.dtype)
7402
  .expand(batch_size, seq_len, -1, -1, -1)
7403
  )
7404
+ # Keep the differentiable StackMemory occupancy/control state in
7405
+ # FP32 even when the live model parameters/activations are BF16.
7406
+ stack_memory_mask = (
7407
+ self.memory_mask.detach()
7408
+ .to(device=hidden_states.device, dtype=torch.float32)
7409
+ .expand(batch_size, seq_len, -1, -1)
7410
  )
7411
  else:
7412
  stack_memory = None
7413
  stack_memory_mask = None
7414
+ all_stack_metrics = () if use_stack_memory else None
7415
 
7416
  position_embeddings = self.rotary_emb(hidden_states, position_ids)
7417
  self.first_layer_fan = (
 
7487
  attn_res_partial = hidden_states # start new block from current output
7488
 
7489
  if use_stack_memory and not use_siamesenorm:
7490
+ hidden_states, stack_memory, stack_memory_mask, stack_layer_metrics = (
7491
  decoder_layer.apply_stack_memory(
7492
+ hidden_states,
7493
+ stack_memory,
7494
+ stack_memory_mask,
7495
+ token_mask=attention_mask,
7496
  )
7497
  )
7498
+ all_stack_metrics = all_stack_metrics + (stack_layer_metrics,)
7499
 
7500
  if use_siamesenorm:
7501
  layer_outputs = decoder_layer.forward_siamesenorm(
 
7504
  position_embeddings=position_embeddings,
7505
  stack_memory=stack_memory,
7506
  stack_memory_mask=stack_memory_mask,
7507
+ stack_token_mask=attention_mask,
7508
  attention_mask=causal_mask,
7509
  first_layer_fan=self.first_layer_fan,
7510
  output_attentions=output_attentions,
 
7520
  if use_stack_memory:
7521
  stack_memory = layer_outputs[2]
7522
  stack_memory_mask = layer_outputs[3]
7523
+ stack_layer_metrics = layer_outputs[4]
7524
+ all_stack_metrics = all_stack_metrics + (stack_layer_metrics,)
7525
+ extras_start = 5
7526
  else:
7527
  extras_start = 2
7528
  else:
 
7583
  if output_hidden_states:
7584
  all_hidden_states = all_hidden_states + (hidden_states,)
7585
 
7586
+ stack_metrics = (
7587
+ torch.stack(all_stack_metrics, dim=0)
7588
+ if all_stack_metrics is not None and len(all_stack_metrics) > 0
7589
+ else None
7590
+ )
7591
+
7592
  if not return_dict:
7593
  base_outputs = tuple(
7594
  v
7595
  for v in [hidden_states, None, all_hidden_states, all_attentions]
7596
  if v is not None
7597
  )
7598
+ if stack_metrics is not None:
7599
+ base_outputs = base_outputs + (stack_metrics,)
7600
  if all_jtokm_aux_stats is not None:
7601
  base_outputs = base_outputs + (all_jtokm_aux_stats,)
7602
  if output_tweo_activations:
 
7609
  hidden_states=all_hidden_states,
7610
  attentions=all_attentions,
7611
  jtokm_aux_stats=all_jtokm_aux_stats,
7612
+ stack_metrics=stack_metrics,
7613
  )
7614
  if output_tweo_activations:
7615
  return outputs, tweo_activations
 
9442
  self._last_nextlat_mse_loss = None
9443
  self._last_nextlat_kl_loss = None
9444
  self._last_jtokm_aux_loss = None
9445
+ self._last_stack_entropy_loss = None
9446
+ self._last_stack_action_entropy = None
9447
+ self._last_stack_action_entropy_normalized = None
9448
+ self._last_stack_push_prob = None
9449
+ self._last_stack_pop_prob = None
9450
+ self._last_stack_noop_prob = None
9451
+ self._last_stack_action_max_prob = None
9452
+ self._last_stack_expected_depth = None
9453
+ self._last_stack_mask_fill_fraction = None
9454
+ self._last_stack_top_occupancy = None
9455
+ self._last_stack_gate_entropy = None
9456
+ self._last_stack_gate_entropy_normalized = None
9457
+ self._last_stack_gate_max_prob = None
9458
+ self._last_stack_gate_effective_slots = None
9459
  self._last_total_loss = None
9460
 
9461
  if config.use_token_generator:
 
9466
  def get_input_embeddings(self):
9467
  return self.model.get_input_embeddings()
9468
 
9469
+ def get_dynamics_metrics(self):
9470
+ """Detached device scalars from the last gradient-enabled training forward.
9471
+
9472
+ Read outside compiled forward at logging cadence. Sampling describes
9473
+ the last microbatch/recomputation, not an average over optimizer steps.
9474
+ This method does not synchronize, transfer activations, or change loss.
9475
+ """
9476
+ if not bool(getattr(self.config, "dynamics_metrics_enabled", False)):
9477
+ return {}
9478
+ metrics = {}
9479
+ generator = getattr(self.model, "token_generator", None)
9480
+ packed = getattr(generator, "_last_dynamics_metrics", None)
9481
+ if packed is not None:
9482
+ metrics.update({f"dynamics/leviathan/{name}": value.detach()
9483
+ for name, value in zip(LEV_NAMES, packed.unbind())})
9484
+ for layer in self.model.layers:
9485
+ module = getattr(layer, "jtok", None)
9486
+ packed = getattr(module, "_last_dynamics_metrics", None)
9487
+ if packed is None:
9488
+ continue
9489
+ kind = "jtokm" if module.use_mixture else "jtok"
9490
+ prefix = f"dynamics/{kind}/layer_{module.layer_idx}/"
9491
+ metrics.update({prefix + name: value.detach() for name, value in
9492
+ unpack_jtok_dynamics(packed, module.num_experts,
9493
+ mixture=module.use_mixture).items()})
9494
+ return metrics
9495
+
9496
+ def get_dynamics_parameter_groups(self):
9497
+ """Only investigated generator/bridge/surface/router parameters, no backbone."""
9498
+ groups = {}
9499
+ generator = getattr(self.model, "token_generator", None)
9500
+ if generator is not None:
9501
+ for name, parameter in generator.named_parameters():
9502
+ if parameter.requires_grad:
9503
+ groups[f"leviathan/{name}"] = (parameter,)
9504
+ for layer in self.model.layers:
9505
+ module = getattr(layer, "jtok", None)
9506
+ if module is not None:
9507
+ kind = "jtokm" if module.use_mixture else "jtok"
9508
+ for name, parameter in module.named_parameters():
9509
+ if parameter.requires_grad:
9510
+ groups[f"{kind}/layer_{module.layer_idx}/{name}"] = (parameter,)
9511
+ return groups
9512
+
9513
  def set_input_embeddings(self, value):
9514
  self.model.set_input_embeddings(value)
9515
 
 
9685
  tweo_activations = None
9686
  hidden_states_tuple = None
9687
  jtokm_aux_stats = None
9688
+ stack_metrics = None
9689
  if isinstance(model_out, tuple):
9690
  outputs = model_out[0]
9691
  if tweo_enabled and len(model_out) >= 2:
 
9702
  if tuple_candidates:
9703
  hidden_states_tuple = tuple_candidates[0]
9704
  for candidate in model_out[1:]:
9705
+ if (
9706
+ bool(getattr(self.config, "use_stack_memory", False))
9707
+ and isinstance(candidate, torch.Tensor)
9708
+ and candidate.ndim == 2
9709
+ and candidate.shape[-1] == 11
9710
+ ):
9711
+ stack_metrics = candidate
9712
  if (
9713
  isinstance(candidate, tuple)
9714
  and candidate
 
9716
  and len(candidate[0]) == 3
9717
  ):
9718
  jtokm_aux_stats = candidate
 
9719
  elif isinstance(outputs, tuple):
9720
  hidden_states = outputs[0]
9721
  if len(outputs) > 2:
 
9724
  hidden_states = outputs.last_hidden_state
9725
  hidden_states_tuple = outputs.hidden_states
9726
  jtokm_aux_stats = getattr(outputs, "jtokm_aux_stats", None)
9727
+ stack_metrics = getattr(outputs, "stack_metrics", None)
9728
  else:
9729
  outputs = model_out
9730
  hidden_states = outputs.last_hidden_state
9731
  hidden_states_tuple = outputs.hidden_states
9732
  jtokm_aux_stats = getattr(outputs, "jtokm_aux_stats", None)
9733
+ stack_metrics = getattr(outputs, "stack_metrics", None)
9734
 
9735
  loss = None
9736
  ntp_loss = None
 
9754
  nextlat_mse_loss = None
9755
  nextlat_kl_loss = None
9756
  jtokm_aux_loss = None
9757
+ stack_entropy_loss = None
9758
  self._last_ntp_loss = None
9759
  self._last_ntp_ce_unweighted = None
9760
  self._last_mile_reweighting_delta = None
 
9781
  self._last_nextlat_mse_loss = None
9782
  self._last_nextlat_kl_loss = None
9783
  self._last_jtokm_aux_loss = None
9784
+ self._last_stack_entropy_loss = None
9785
+ self._last_stack_action_entropy = None
9786
+ self._last_stack_action_entropy_normalized = None
9787
+ self._last_stack_push_prob = None
9788
+ self._last_stack_pop_prob = None
9789
+ self._last_stack_noop_prob = None
9790
+ self._last_stack_action_max_prob = None
9791
+ self._last_stack_expected_depth = None
9792
+ self._last_stack_mask_fill_fraction = None
9793
+ self._last_stack_top_occupancy = None
9794
+ self._last_stack_gate_entropy = None
9795
+ self._last_stack_gate_entropy_normalized = None
9796
+ self._last_stack_gate_max_prob = None
9797
+ self._last_stack_gate_effective_slots = None
9798
  self._last_total_loss = None
9799
  if labels is not None:
9800
  if self.ntp_loss_backend == "liger":
 
9850
  ntp_loss = ntp_result
9851
  loss = ntp_loss
9852
 
9853
+ # StackTrans entropy regularizer. The backbone returns one compact
9854
+ # metric row per decoder layer; column 0 is the differentiable mean
9855
+ # entropy H(push,pop,noop). Averaging across layers keeps lambda
9856
+ # invariant to model depth, batch size, sequence length, and head count.
9857
+ if bool(getattr(self.config, "use_stack_memory", False)):
9858
+ if stack_metrics is None or stack_metrics.numel() == 0:
9859
+ raise RuntimeError(
9860
+ "StackMemory is active but no stack metrics were returned."
9861
+ )
9862
+ stack_action_entropy = stack_metrics[:, 0].mean()
9863
+ stack_entropy_weight = torch.as_tensor(
9864
+ float(getattr(self.config, "stack_entropy_loss_weight", 1e-3)),
9865
+ device=stack_action_entropy.device,
9866
+ dtype=torch.float32,
9867
+ )
9868
+ stack_entropy_loss = stack_entropy_weight * stack_action_entropy
9869
+ loss = loss + stack_entropy_loss
9870
+
9871
  if tweo_enabled:
9872
  if not tweo_activations:
9873
  raise ValueError(
 
10153
  self._last_jtokm_aux_loss = (
10154
  jtokm_aux_loss.detach() if jtokm_aux_loss is not None else None
10155
  )
10156
+ if bool(getattr(self.config, "use_stack_memory", False)):
10157
+ detached_stack = stack_metrics.detach().float()
10158
+ action_entropy = detached_stack[:, 0].mean()
10159
+ gate_entropy = detached_stack[:, 8].mean()
10160
+ self._last_stack_entropy_loss = (
10161
+ stack_entropy_loss.detach()
10162
+ if stack_entropy_loss is not None
10163
+ else action_entropy.new_zeros(())
10164
+ )
10165
+ self._last_stack_action_entropy = action_entropy
10166
+ self._last_stack_action_entropy_normalized = action_entropy / math.log(3.0)
10167
+ self._last_stack_push_prob = detached_stack[:, 1].mean()
10168
+ self._last_stack_pop_prob = detached_stack[:, 2].mean()
10169
+ self._last_stack_noop_prob = detached_stack[:, 3].mean()
10170
+ self._last_stack_action_max_prob = detached_stack[:, 4].mean()
10171
+ self._last_stack_expected_depth = detached_stack[:, 5].mean()
10172
+ self._last_stack_mask_fill_fraction = detached_stack[:, 6].mean()
10173
+ self._last_stack_top_occupancy = detached_stack[:, 7].mean()
10174
+ self._last_stack_gate_entropy = gate_entropy
10175
+ self._last_stack_gate_entropy_normalized = gate_entropy / math.log(
10176
+ float(self.config.stack_slots)
10177
+ )
10178
+ self._last_stack_gate_max_prob = detached_stack[:, 9].mean()
10179
+ self._last_stack_gate_effective_slots = detached_stack[:, 10].mean()
10180
  self._last_total_loss = loss.detach() if loss is not None else None
10181
  logits = None
10182
  else: