Tessera-135M-Gate / modeling.py
yava-code's picture
Publish Tessera-1B-Nano checkpoint
7709bf3 verified
Raw History Blame Contribute Delete
11.6 kB
from __future__ import annotations
from dataclasses import dataclass
from typing import Any, Literal
import torch
import torch.nn.functional as F
from torch import Tensor, nn
from transformers import AutoModelForCausalLM, GenerationMixin, LlamaConfig, PreTrainedModel
from transformers.modeling_outputs import ModelOutput
from transformers.models.llama.modeling_llama import (
LlamaDecoderLayer,
LlamaForCausalLM,
LlamaRotaryEmbedding,
)
from .configuration import NcpSmolConfig
from .quantizer import ProductVectorQuantizer
ConceptMode = Literal["predicted", "zero", "shuffle"]
@dataclass
class NcpCausalLMOutput(ModelOutput):
loss: Tensor | None = None
logits: Tensor | None = None
ntp_loss: Tensor | None = None
ncp_loss: Tensor | None = None
vq_loss: Tensor | None = None
code_indices: Tensor | None = None
class NcpSmolForCausalLM(PreTrainedModel, GenerationMixin):
config_class = NcpSmolConfig
base_model_prefix = "backbone"
main_input_name = "input_ids"
_supports_cache_class = False
_no_split_modules = ["LlamaDecoderLayer"]
_tied_weights_keys = ["backbone.lm_head.weight"]
def __init__(self, config: NcpSmolConfig) -> None:
super().__init__(config)
backbone_config = LlamaConfig.from_dict(config.backbone_config)
self.backbone = LlamaForCausalLM(backbone_config)
concept_config = LlamaConfig.from_dict(config.backbone_config)
concept_config._attn_implementation = "eager"
self.concept_layers = nn.ModuleList(
[LlamaDecoderLayer(concept_config, i) for i in range(config.concept_layers)]
)
self.concept_rotary = LlamaRotaryEmbedding(concept_config)
self.quantizer = ProductVectorQuantizer(
hidden_size=config.hidden_size,
segments=config.segments,
codebook_size=config.codebook_size,
)
self.concept_head = nn.Linear(
config.hidden_size,
config.segments * config.codebook_size,
)
self._hook_state: dict[str, Any] = {}
self._hook = self.backbone.model.layers[config.insert_layer].register_forward_pre_hook(
self._inject_concepts,
with_kwargs=True,
)
self.post_init()
self.init_concept_from_backbone()
self.tie_weights()
@classmethod
def from_backbone(
cls,
model_name_or_path: str,
*,
revision: str | None = None,
chunk_size: int = 4,
segments: int | None = None,
codebook_size: int = 64,
concept_layers: int = 2,
insert_layer: int = 1,
ncp_target: str = "continuous",
ncp_weight: float = 1.0,
vq_weight: float = 1.0,
**load_kwargs: Any,
) -> NcpSmolForCausalLM:
base = AutoModelForCausalLM.from_pretrained(
model_name_or_path,
revision=revision,
**load_kwargs,
)
if not isinstance(base, LlamaForCausalLM):
raise TypeError("ncp-smol currently requires a Llama-compatible backbone")
config = NcpSmolConfig(
backbone_config=base.config.to_dict(),
base_model_name_or_path=model_name_or_path,
base_model_revision=revision,
chunk_size=chunk_size,
segments=segments,
codebook_size=codebook_size,
concept_layers=concept_layers,
insert_layer=insert_layer,
ncp_target=ncp_target,
ncp_weight=ncp_weight,
vq_weight=vq_weight,
)
model = cls(config)
model.backbone.load_state_dict(base.state_dict())
model.init_concept_from_backbone()
model.tie_weights()
return model
def init_concept_from_backbone(self) -> None:
source = self.backbone.model.layers
for index, layer in enumerate(self.concept_layers):
layer.load_state_dict(source[index % len(source)].state_dict())
def get_input_embeddings(self) -> nn.Module:
return self.backbone.get_input_embeddings()
def set_input_embeddings(self, value: nn.Module) -> None:
self.backbone.set_input_embeddings(value)
def get_output_embeddings(self) -> nn.Module:
return self.backbone.get_output_embeddings()
def set_output_embeddings(self, value: nn.Module) -> None:
self.backbone.set_output_embeddings(value)
def tie_weights(self) -> None:
self.backbone.tie_weights()
def _concept_mask(self, valid: Tensor, dtype: torch.dtype) -> Tensor:
batch, length = valid.shape
mask = torch.full(
(length, length),
torch.finfo(dtype).min,
dtype=dtype,
device=valid.device,
).triu(diagonal=1)
mask = mask.view(1, 1, length, length).expand(batch, 1, length, length).clone()
return mask.masked_fill(~valid[:, None, None, :], torch.finfo(dtype).min)
def _concept_path(
self,
hidden: Tensor,
attention_mask: Tensor | None,
mode: ConceptMode,
) -> tuple[Tensor, dict[str, Tensor]]:
batch, length, width = hidden.shape
chunks = length // self.config.chunk_size
zero = hidden.sum() * 0
if chunks == 0:
return torch.zeros_like(hidden), {
"ncp_loss": zero,
"vq_loss": zero,
"code_indices": torch.empty(
batch, 0, self.config.segments, dtype=torch.long, device=hidden.device
),
}
usable = chunks * self.config.chunk_size
concepts = (
hidden[:, :usable]
.reshape(
batch,
chunks,
self.config.chunk_size,
width,
)
.mean(dim=2)
)
if attention_mask is None:
valid = torch.ones(batch, chunks, dtype=torch.bool, device=hidden.device)
else:
valid = (
attention_mask[:, :usable]
.bool()
.reshape(batch, chunks, self.config.chunk_size)
.all(dim=2)
)
quantized = self.quantizer(concepts)
positions = torch.arange(chunks, device=hidden.device)
position_ids = (positions * self.config.chunk_size).unsqueeze(0).expand(batch, chunks)
concept_hidden = concepts
rotary = self.concept_rotary(concept_hidden, position_ids)
causal_mask = self._concept_mask(valid, concept_hidden.dtype)
for layer in self.concept_layers:
concept_hidden = layer(
concept_hidden,
attention_mask=causal_mask,
position_ids=position_ids,
cache_position=positions,
position_embeddings=rotary,
use_cache=False,
)
logits = self.concept_head(concept_hidden).view(
batch,
chunks,
self.config.segments,
self.config.codebook_size,
)
predicted = self.quantizer.expected(logits)
if chunks > 1 and valid[:, 1:].any():
target_mask = valid[:, 1:].unsqueeze(-1)
target = quantized.codes if self.config.ncp_target == "quantized" else concepts
squared = (predicted[:, :-1] - target[:, 1:].detach()).square()
ncp_loss = (squared * target_mask).sum() / (target_mask.sum().clamp_min(1) * width)
else:
ncp_loss = zero
if mode == "zero":
predicted = torch.zeros_like(predicted)
elif mode == "shuffle":
predicted = predicted.roll(shifts=1, dims=0)
feedback = self._align_feedback(predicted.detach(), length)
return feedback, {
"ncp_loss": ncp_loss,
"vq_loss": quantized.loss,
"code_indices": quantized.indices,
}
def _align_feedback(self, predicted: Tensor, length: int) -> Tensor:
batch, _, width = predicted.shape
prefix = torch.zeros(batch, 1, width, dtype=predicted.dtype, device=predicted.device)
repeated = torch.cat((prefix, predicted), dim=1).repeat_interleave(
self.config.chunk_size,
dim=1,
)
return repeated[:, 1 : length + 1]
def _inject_concepts(
self,
_module: nn.Module,
args: tuple[Any, ...],
kwargs: dict[str, Any],
) -> tuple[tuple[Any, ...], dict[str, Any]]:
if not self._hook_state.get("active", False):
return args, kwargs
hidden = args[0] if args else kwargs["hidden_states"]
feedback, aux = self._concept_path(
hidden,
self._hook_state.get("attention_mask"),
self._hook_state["mode"],
)
self._hook_state["aux"] = aux
if args:
args = (hidden + feedback, *args[1:])
else:
kwargs["hidden_states"] = hidden + feedback
return args, kwargs
def forward(
self,
input_ids: Tensor | None = None,
attention_mask: Tensor | None = None,
labels: Tensor | None = None,
inputs_embeds: Tensor | None = None,
concept_mode: ConceptMode = "predicted",
use_cache: bool | None = None,
return_dict: bool = True,
**kwargs: Any,
) -> NcpCausalLMOutput | tuple[Tensor, ...]:
if concept_mode not in {"predicted", "zero", "shuffle"}:
raise ValueError(f"unknown concept_mode: {concept_mode}")
self._hook_state = {
"active": True,
"attention_mask": attention_mask,
"mode": concept_mode,
"aux": None,
}
try:
outputs = self.backbone(
input_ids=input_ids,
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
labels=None,
use_cache=False,
return_dict=True,
**kwargs,
)
aux = self._hook_state["aux"]
finally:
self._hook_state["active"] = False
if aux is None:
raise RuntimeError("concept injection hook did not run")
ntp_loss = None
total_loss = None
if labels is not None:
shift_logits = outputs.logits[:, :-1].float().contiguous()
shift_labels = labels[:, 1:].contiguous()
ntp_loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1),
ignore_index=-100,
)
total_loss = (
ntp_loss
+ self.config.ncp_weight * aux["ncp_loss"]
+ self.config.vq_weight * aux["vq_loss"]
)
result = NcpCausalLMOutput(
loss=total_loss,
logits=outputs.logits,
ntp_loss=ntp_loss,
ncp_loss=aux["ncp_loss"],
vq_loss=aux["vq_loss"],
code_indices=aux["code_indices"],
)
if return_dict:
return result
return tuple(value for value in result.values() if value is not None)
def prepare_inputs_for_generation(
self,
input_ids: Tensor,
attention_mask: Tensor | None = None,
**_: Any,
) -> dict[str, Any]:
return {
"input_ids": input_ids,
"attention_mask": attention_mask,
"use_cache": False,
}
NcpSmolForCausalLM.register_for_auto_class("AutoModelForCausalLM")