From 32a6a5894576473696993a6f0ef4493bb21ecaf9 Mon Sep 17 00:00:00 2001 From: Divyam Talwar Date: Mon, 24 Aug 2026 00:29:31 +0530 Subject: [PATCH] feat: expose opt-in per-layer geometric diagnostics --- README.md | 25 +++++++++ model/ouroboros.py | 105 +++++++++++++++++++++++++++++--------- tests/test_diagnostics.py | 74 +++++++++++++++++++++++++++ 3 files changed, 181 insertions(+), 23 deletions(-) create mode 100644 tests/test_diagnostics.py diff --git a/README.md b/README.md index b01d4d7..7d53e6f 100644 --- a/README.md +++ b/README.md @@ -61,6 +61,31 @@ print(logits.shape, float(loss)) `hidden_size` must equal `num_attention_heads * head_dim`. Inputs longer than `block_size` fail explicitly. `crop_block_size(n)` only decreases that limit. +## Opt-in residual diagnostics + +Set `ouroboros_collect_diagnostics=True` in the config to collect per-layer +metrics during a forward pass: + +```python +config = OuroborosConfig( + vocab_size=256, + hidden_size=64, + num_hidden_layers=2, + num_attention_heads=4, + head_dim=16, + ouroboros_collect_diagnostics=True, +) +model = OuroborosModel(config) +model(torch.randint(0, 256, (1, 16))) +print(model.get_residual_diagnostics()) +``` + +Each attention and MLP residual reports beta range/mean, absolute gain along the +updated axis `|1-beta|`, projection error before/after the write, and update +norm. Collection is disabled by default because scalar extraction synchronizes +accelerators. Metrics describe the realized update; they do not establish a +quality improvement. + ## Allocation-aware update The default geometric update uses `torch.addcmul(x, k, delta)` rather than diff --git a/model/ouroboros.py b/model/ouroboros.py index e293481..d451233 100644 --- a/model/ouroboros.py +++ b/model/ouroboros.py @@ -118,9 +118,9 @@ def __init__(self, config): self.k_eps = float(getattr(config, "ouroboros_k_eps", 1e-6)) self.v_sigmoid = bool(getattr(config, "ouroboros_v_sigmoid", True)) - self.v_sigmoid_scale: float = float(getattr(config, "ouroboros_v_sigmoid_scale", 4.0)) + self.v_sigmoid_scale = float(getattr(config, "ouroboros_v_sigmoid_scale", 4.0)) self.v_constant = bool(getattr(config, "ouroboros_v_constant", False)) - self.v_constant_value: float = float(getattr(config, "ouroboros_v_constant_value", 2.0)) + self.v_constant_value = float(getattr(config, "ouroboros_v_constant_value", 2.0)) self.beta_single_linear = bool(getattr(config, "ouroboros_beta_single_linear", True)) if self.beta_single_linear: @@ -129,42 +129,72 @@ def __init__(self, config): beta_hidden_size = int(getattr(config, "ouroboros_beta_hidden_size", 128)) if beta_hidden_size <= 0: raise ValueError("ouroboros_beta_hidden_size must be positive.") - self.beta_in = nn.Linear(hidden_size, beta_hidden_size, bias=False) self.beta_out = nn.Linear(beta_hidden_size, 1, bias=True) - # v is a scalar in the d_v=1 regime; project sublayer output to a scalar. self.v_proj = nn.Linear(hidden_size, 1, bias=True) - beta_init = float(getattr(config, "ouroboros_beta_init", 0.0)) - beta_init = min(max(beta_init, 0.0), 2.0) - beta_init_p = beta_init / 2.0 + beta_init = min(max(float(getattr(config, "ouroboros_beta_init", 0.0)), 0.0), 2.0) with torch.no_grad(): + beta_bias = _logit(beta_init / 2.0) if self.beta_single_linear: - self.beta.bias.fill_(_logit(beta_init_p)) + self.beta.bias.fill_(beta_bias) else: - self.beta_out.bias.fill_(_logit(beta_init_p)) + self.beta_out.bias.fill_(beta_bias) - def forward(self, x: torch.Tensor, *, k_in: torch.Tensor, context: torch.Tensor) -> torch.Tensor: + def _components(self, x, k_in, context): k = F.normalize(k_in, p=2, dim=-1, eps=self.k_eps) - if self.beta_single_linear: beta_logits = self.beta(context).float() else: beta_logits = self.beta_out(torch.tanh(self.beta_in(context))).float() - beta = 2.0 * torch.sigmoid(beta_logits) # fp32 - - proj = torch.sum(k * x, dim=-1, keepdim=True, dtype=torch.float32) # fp32 + beta = 2.0 * torch.sigmoid(beta_logits) + projection = torch.sum(k * x, dim=-1, keepdim=True, dtype=torch.float32) if self.v_constant: - v = torch.full_like(proj, self.v_constant_value) + target = torch.full_like(projection, self.v_constant_value) else: - v = self.v_proj(x) + target = self.v_proj(x) if self.v_sigmoid: - v = torch.sigmoid(v) * self.v_sigmoid_scale + target = torch.sigmoid(target) * self.v_sigmoid_scale + delta = (beta * (target - projection)).to(dtype=x.dtype) + return k, beta, target, projection, delta - delta = (beta * (v - proj)).to(dtype=x.dtype) # (B, T, 1) - return geometric_update(x, k, delta) + @staticmethod + def _diagnostics(x, output, k, beta, target, projection): + with torch.no_grad(): + post_projection = torch.sum( + k * output, dim=-1, keepdim=True, dtype=torch.float32 + ) + return { + "beta_mean": float(beta.mean().item()), + "beta_min": float(beta.min().item()), + "beta_max": float(beta.max().item()), + "axis_gain_abs_max": float((1.0 - beta).abs().max().item()), + "projection_error_before_mean": float( + (target.float() - projection).abs().mean().item() + ), + "projection_error_after_mean": float( + (target.float() - post_projection).abs().mean().item() + ), + "update_norm_mean": float( + (output.float() - x.float()).norm(dim=-1).mean().item() + ), + } + + def forward( + self, + x: torch.Tensor, + *, + k_in: torch.Tensor, + context: torch.Tensor, + return_diagnostics: bool = False, + ): + k, beta, target, projection, delta = self._components(x, k_in, context) + output = geometric_update(x, k, delta) + if not return_diagnostics: + return output + return output, self._diagnostics(x, output, k, beta, target, projection) class OuroborosBlock(nn.Module): @@ -174,19 +204,35 @@ def __init__(self, config): self.mlp = MLP(config) self.ouroboros_attn = OuroborosResidual(config) self.ouroboros_mlp = OuroborosResidual(config) - # Define RMSNorm layers once in the module self.ln_1 = RMSNorm(config.hidden_size) self.ln_2 = RMSNorm(config.hidden_size) + self.collect_diagnostics = bool( + getattr(config, "ouroboros_collect_diagnostics", False) + ) + self.last_diagnostics = {} def forward(self, x): - # Apply pre-norm before sublayers x_norm = self.ln_1(x) k_attn = self.attn(x_norm) - x = self.ouroboros_attn(x, k_in=k_attn, context=x_norm) + if self.collect_diagnostics: + x, attn_diagnostics = self.ouroboros_attn( + x, k_in=k_attn, context=x_norm, return_diagnostics=True + ) + else: + x = self.ouroboros_attn(x, k_in=k_attn, context=x_norm) x_norm = self.ln_2(x) k_mlp = self.mlp(x_norm) - x = self.ouroboros_mlp(x, k_in=k_mlp, context=x_norm) + if self.collect_diagnostics: + x, mlp_diagnostics = self.ouroboros_mlp( + x, k_in=k_mlp, context=x_norm, return_diagnostics=True + ) + self.last_diagnostics = { + "attention": attn_diagnostics, + "mlp": mlp_diagnostics, + } + else: + x = self.ouroboros_mlp(x, k_in=k_mlp, context=x_norm) return x @dataclass @@ -218,6 +264,7 @@ class OuroborosConfig(PretrainedConfig): ouroboros_v_sigmoid_scale: float = 4.0 ouroboros_v_constant: bool = False ouroboros_v_constant_value: float = 2.0 + ouroboros_collect_diagnostics: bool = False # Initialize beta; clamped to [0, 2]. Use 1.0 by default for baseline comparability. ouroboros_beta_init: float = 1.0 @@ -286,6 +333,18 @@ def forward(self, idx, targets=None, return_logits=True, output_all_seq=False): return logits, loss + def get_residual_diagnostics(self): + """Return per-layer metrics from the most recent instrumented forward pass.""" + + if not getattr(self.config, "ouroboros_collect_diagnostics", False): + raise RuntimeError( + "set ouroboros_collect_diagnostics=True before model construction" + ) + return [ + {"layer": index, **block.last_diagnostics} + for index, block in enumerate(self.transformer.h) + ] + def crop_block_size(self, block_size): """Decrease the configured maximum sequence length.""" block_size = int(block_size) diff --git a/tests/test_diagnostics.py b/tests/test_diagnostics.py new file mode 100644 index 0000000..fd67a9c --- /dev/null +++ b/tests/test_diagnostics.py @@ -0,0 +1,74 @@ +import json + +import torch + +from model import OuroborosConfig, OuroborosModel, OuroborosResidual + + +def tiny_config(**overrides): + values = dict( + vocab_size=32, + num_hidden_layers=2, + num_attention_heads=2, + hidden_size=16, + head_dim=8, + block_size=16, + using_groupnorm=False, + ) + values.update(overrides) + return OuroborosConfig(**values) + + +def test_beta_one_reaches_constant_projection_target(): + residual = OuroborosResidual( + tiny_config( + ouroboros_beta_init=1.0, + ouroboros_v_constant=True, + ouroboros_v_constant_value=2.0, + ) + ) + x = torch.randn(2, 5, 16) + k = torch.randn_like(x) + context = torch.zeros_like(x) + + _, diagnostics = residual( + x, k_in=k, context=context, return_diagnostics=True + ) + + assert diagnostics["beta_min"] == 1.0 + assert diagnostics["beta_max"] == 1.0 + assert diagnostics["projection_error_after_mean"] < 1e-5 + assert diagnostics["projection_error_after_mean"] < diagnostics[ + "projection_error_before_mean" + ] + + +def test_model_collects_json_serializable_layer_diagnostics_without_output_drift(): + torch.manual_seed(3) + baseline = OuroborosModel(tiny_config(ouroboros_collect_diagnostics=False)).eval() + instrumented = OuroborosModel(tiny_config(ouroboros_collect_diagnostics=True)).eval() + instrumented.load_state_dict(baseline.state_dict()) + tokens = torch.randint(0, 32, (2, 7)) + + expected, _ = baseline(tokens, output_all_seq=True) + actual, _ = instrumented(tokens, output_all_seq=True) + torch.testing.assert_close(actual, expected) + + report = instrumented.get_residual_diagnostics() + assert len(report) == 2 + for layer in report: + for name in ("attention", "mlp"): + metrics = layer[name] + assert 0.0 <= metrics["beta_min"] <= metrics["beta_max"] <= 2.0 + assert 0.0 <= metrics["axis_gain_abs_max"] <= 1.0 + json.dumps(report) + + +def test_diagnostics_fail_clearly_when_disabled(): + model = OuroborosModel(tiny_config()) + try: + model.get_residual_diagnostics() + except RuntimeError as error: + assert "ouroboros_collect_diagnostics=True" in str(error) + else: + raise AssertionError("disabled diagnostics should fail clearly")