nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
7 kB
"""Norm-bounded algebraic graph reads around an otherwise frozen LM."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from torch import nn
from torch.nn import functional as F
from strata.modeling.algebra import BoundedRelativeReZero
@dataclass(frozen=True, slots=True)
class ComposeLMOutput:
logits: torch.Tensor
loss: torch.Tensor | None
hidden_states: torch.Tensor
graph_update: torch.Tensor
relative_update_rms: torch.Tensor
@dataclass(frozen=True, slots=True)
class SparseAlgebraGraphRead:
"""CSR-like graph values addressed to selected batch/token positions."""
values: torch.Tensor
batch_indices: torch.Tensor
token_indices: torch.Tensor
reliability: torch.Tensor
def validate(self, hidden: torch.Tensor, graph_dim: int) -> None:
count = self.values.shape[0]
if self.values.ndim != 2 or self.values.shape[1] != graph_dim:
raise ValueError("sparse graph values must have shape [reads, graph_dim]")
if self.batch_indices.shape != (count,) or self.token_indices.shape != (count,):
raise ValueError("sparse graph indices must have shape [reads]")
if self.reliability.shape != (count,):
raise ValueError("sparse graph reliability must have shape [reads]")
if count and (
int(self.batch_indices.min()) < 0
or int(self.batch_indices.max()) >= hidden.shape[0]
or int(self.token_indices.min()) < 0
or int(self.token_indices.max()) >= hidden.shape[1]
):
raise ValueError("sparse graph index is out of range")
class AlgebraGraphReadAdapter(nn.Module):
"""Project exact graph values and inject them through the fixed BRR bound."""
def __init__(self, graph_dim: int, hidden_dim: int, *, gamma_max: float = 0.02) -> None:
super().__init__()
self.graph_projection = nn.Linear(graph_dim, hidden_dim, bias=False)
self.brr = BoundedRelativeReZero(hidden_dim, gamma_max=gamma_max, per_channel=False)
def forward(
self,
hidden: torch.Tensor,
graph_values: torch.Tensor,
reliability: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
if graph_values.shape[:-1] != hidden.shape[:-1]:
raise ValueError("graph values must align with token hidden states")
if reliability.shape != hidden.shape[:-1]:
raise ValueError("reliability must have shape [batch, sequence]")
active = reliability != 0
if not bool(active.any()):
return hidden, torch.zeros_like(hidden), hidden.new_zeros(hidden.shape[:-1])
projection_dtype = self.graph_projection.weight.dtype
if bool(active.all()):
read = self.graph_projection(graph_values.to(projection_dtype)).to(hidden.dtype)
output = self.brr(hidden, read, reliability)
relative = (output.update_rms / output.hidden_rms.clamp_min(1e-8)).squeeze(-1)
return output.hidden, output.update, relative
active_hidden = hidden[active]
active_read = self.graph_projection(graph_values[active].to(projection_dtype)).to(hidden.dtype)
active_output = self.brr(active_hidden, active_read, reliability[active])
updated = hidden.clone()
updated[active] = active_output.hidden
update = torch.zeros_like(hidden)
update[active] = active_output.update
relative = hidden.new_zeros(hidden.shape[:-1])
relative[active] = (
active_output.update_rms / active_output.hidden_rms.clamp_min(1e-8)
).squeeze(-1)
return updated, update, relative
def forward_sparse(
self,
hidden: torch.Tensor,
read: SparseAlgebraGraphRead,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
read.validate(hidden, self.graph_projection.in_features)
if read.values.shape[0] == 0:
return hidden, torch.zeros_like(hidden), hidden.new_zeros(hidden.shape[:-1])
batch = read.batch_indices.to(device=hidden.device, dtype=torch.long)
token = read.token_indices.to(device=hidden.device, dtype=torch.long)
selected_hidden = hidden[batch, token]
projection_dtype = self.graph_projection.weight.dtype
selected_read = self.graph_projection(
read.values.to(device=hidden.device, dtype=projection_dtype)
).to(hidden.dtype)
selected_output = self.brr(
selected_hidden,
selected_read,
read.reliability.to(device=hidden.device, dtype=hidden.dtype),
)
updated = hidden.clone()
updated[batch, token] = selected_output.hidden
update = torch.zeros_like(hidden)
update[batch, token] = selected_output.update
relative = hidden.new_zeros(hidden.shape[:-1])
relative[batch, token] = (
selected_output.update_rms / selected_output.hidden_rms.clamp_min(1e-8)
).squeeze(-1)
return updated, update, relative
class FrozenLMWithAlgebraRead(nn.Module):
"""Attach graph reads after a frozen decoder without changing its base path.
The wrapped model must expose ``forward_hidden``, ``final_norm``, and
``lm_head``. When graph reads are disabled, the adapter is not called.
"""
def __init__(self, base_model: nn.Module, adapter: AlgebraGraphReadAdapter) -> None:
super().__init__()
self.base_model = base_model
self.adapter = adapter
for parameter in self.base_model.parameters():
parameter.requires_grad_(False)
def forward(
self,
input_ids: torch.Tensor,
*,
attention_mask: torch.Tensor | None = None,
labels: torch.Tensor | None = None,
graph_values: torch.Tensor | None = None,
graph_reliability: torch.Tensor | None = None,
graph_enabled: bool = False,
) -> ComposeLMOutput:
hidden, _mask = self.base_model.forward_hidden(input_ids, attention_mask=attention_mask)
update = torch.zeros_like(hidden)
relative = hidden.new_zeros(hidden.shape[:-1])
if graph_enabled:
if graph_values is None or graph_reliability is None:
raise ValueError("enabled graph reads require values and reliability")
hidden, update, relative = self.adapter(hidden, graph_values, graph_reliability)
normalized = self.base_model.final_norm(hidden)
logits = self.base_model.lm_head(normalized)
loss = None
if labels is not None:
loss = F.cross_entropy(
logits[:, :-1].contiguous().view(-1, logits.shape[-1]),
labels[:, 1:].contiguous().view(-1),
ignore_index=-100,
)
return ComposeLMOutput(logits, loss, normalized, update, relative)
__all__ = [
"AlgebraGraphReadAdapter",
"ComposeLMOutput",
"FrozenLMWithAlgebraRead",
"SparseAlgebraGraphRead",
]