Update modeling_neollm.py
Browse files- 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 |
-
|
| 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 |
-
|
| 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 |
-
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 *
|
| 1319 |
-
new_mask = (masks *
|
| 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.
|
|
|
|
|
|
|
| 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 |
-
|
| 1342 |
-
#
|
| 1343 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1344 |
|
| 1345 |
-
memory_output = (
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
|
|
|
|
|
|
| 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 =
|
| 5652 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 6824 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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,
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
| 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:
|