Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 25 additions & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
105 changes: 82 additions & 23 deletions model/ouroboros.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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):
Expand All @@ -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
Expand Down Expand Up @@ -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

Expand Down Expand Up @@ -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)
Expand Down
74 changes: 74 additions & 0 deletions tests/test_diagnostics.py
Original file line number Diff line number Diff line change
@@ -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")
Loading