dot / dot_rd /model.py
MTEnt's picture
Add files using upload-large-folder tool
954544e verified
Raw History Blame Contribute Delete
10.5 kB
from __future__ import annotations
import copy
import math
from dataclasses import asdict
from typing import Any
import torch
from torch import nn
from transformers import AutoModelForCausalLM
from transformers.masking_utils import create_causal_mask, create_recurrent_attention_mask
from transformers.modeling_outputs import CausalLMOutputWithPast
from .config import ModelConfig
class RecurrentDepthCore(nn.Module):
"""A weight-tied group of native decoder layers with identity-preserving gates."""
def __init__(self, source_layers: list[nn.Module], max_loops: int, active_loops: int) -> None:
super().__init__()
if not source_layers:
raise ValueError("source_layers must not be empty")
self.layers = nn.ModuleList(copy.deepcopy(source_layers))
self.max_loops = max_loops
self.active_loops = active_loops
self.loop_scale_logits = nn.Parameter(torch.zeros(max_loops, dtype=torch.float32))
@property
def loop_scales(self) -> torch.Tensor:
return torch.tanh(self.loop_scale_logits)
def set_active_loops(self, loops: int) -> None:
if not 1 <= loops <= self.max_loops:
raise ValueError(f"loops must be within [1, {self.max_loops}]")
self.active_loops = loops
def set_initial_scale(self, scale: float) -> None:
if not 0.0 <= scale < 1.0:
raise ValueError("scale must be within [0, 1)")
raw = 0.0 if scale == 0.0 else math.atanh(scale)
with torch.no_grad():
self.loop_scale_logits.zero_()
self.loop_scale_logits[: self.active_loops].fill_(raw)
def open_inactive_zero_scales(self, scale: float) -> None:
"""Open newly activated loops without changing gates learned in earlier stages."""
if not 0.0 <= scale < 1.0:
raise ValueError("scale must be within [0, 1)")
if scale == 0.0:
return
raw = math.atanh(scale)
with torch.no_grad():
active = self.loop_scale_logits[: self.active_loops]
active[active == 0.0] = raw
def forward(
self,
hidden_states: torch.Tensor,
*,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
masks: dict[str, torch.Tensor | None],
position_ids: torch.LongTensor | None,
**kwargs: Any,
) -> torch.Tensor:
for loop_index in range(self.active_loops):
candidate = hidden_states
for layer in self.layers:
candidate = layer(
candidate,
position_embeddings=position_embeddings,
attention_mask=masks[layer.block_type],
position_ids=position_ids,
past_key_values=None,
use_cache=False,
**kwargs,
)
scale = self.loop_scales[loop_index].to(device=hidden_states.device, dtype=hidden_states.dtype)
hidden_states = hidden_states + scale * (candidate - hidden_states)
return hidden_states
class DotRecurrentDepthModel(nn.Module):
"""Qwen3.5 causal LM with a trainable recurrent block inserted at mid-depth.
Cache-backed decoding is deliberately rejected in v0.1. Reusing native layer cache
indices across recurrent passes would silently corrupt state. Training and correctness
evaluation run with ``use_cache=False`` until a dedicated recurrent cache exists.
"""
def __init__(self, backbone: nn.Module, model_config: ModelConfig) -> None:
super().__init__()
self.backbone = backbone
self.model_config = model_config
layers = self.backbone.model.layers
layer_count = len(layers)
if model_config.insertion_after >= layer_count - 1:
raise ValueError(
f"insertion_after={model_config.insertion_after} must leave at least one trailing layer"
)
if max(model_config.source_layers) >= layer_count:
raise ValueError(
f"source layer {max(model_config.source_layers)} exceeds backbone layer count {layer_count}"
)
source_layers = [layers[index] for index in model_config.source_layers]
self.reasoning_core = RecurrentDepthCore(
source_layers=source_layers,
max_loops=model_config.max_loops,
active_loops=model_config.active_loops,
)
@classmethod
def from_pretrained(
cls,
model_config: ModelConfig,
*,
dtype: torch.dtype = torch.bfloat16,
device_map: str | dict[str, Any] | None = None,
) -> "DotRecurrentDepthModel":
backbone = AutoModelForCausalLM.from_pretrained(
model_config.base_model,
dtype=dtype,
attn_implementation=model_config.attention_implementation,
device_map=device_map,
low_cpu_mem_usage=True,
)
return cls(backbone, model_config)
@property
def config(self) -> Any:
return self.backbone.config
def freeze_backbone(self) -> None:
self.backbone.requires_grad_(False)
self.reasoning_core.requires_grad_(True)
def trainable_parameter_count(self) -> int:
return sum(parameter.numel() for parameter in self.parameters() if parameter.requires_grad)
def total_parameter_count(self) -> int:
return sum(parameter.numel() for parameter in self.parameters())
def enable_gradient_checkpointing(self) -> None:
self.backbone.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False}
)
checkpoint_function = None
for layer in self.backbone.model.layers:
checkpoint_function = getattr(layer, "_gradient_checkpointing_func", None)
if checkpoint_function is not None:
break
if checkpoint_function is None:
raise RuntimeError("backbone did not install a gradient checkpoint function")
for layer in self.reasoning_core.layers:
layer.gradient_checkpointing = True
layer._gradient_checkpointing_func = checkpoint_function
def forward(
self,
input_ids: torch.LongTensor | None = None,
attention_mask: torch.Tensor | None = None,
position_ids: torch.LongTensor | None = None,
inputs_embeds: torch.FloatTensor | None = None,
labels: torch.LongTensor | None = None,
use_cache: bool | None = False,
logits_to_keep: int | torch.Tensor = 0,
**kwargs: Any,
) -> CausalLMOutputWithPast:
if use_cache:
raise ValueError("recurrent-depth v0.1 does not support cache-backed decoding")
if (input_ids is None) == (inputs_embeds is None):
raise ValueError("specify exactly one of input_ids or inputs_embeds")
text_model = self.backbone.model
if inputs_embeds is None:
inputs_embeds = text_model.embed_tokens(input_ids)
if position_ids is None:
sequence_positions = torch.arange(inputs_embeds.shape[1], device=inputs_embeds.device)
position_ids = sequence_positions.view(1, 1, -1).expand(4, inputs_embeds.shape[0], -1)
elif position_ids.ndim == 2:
position_ids = position_ids[None, ...].expand(4, position_ids.shape[0], -1)
if position_ids.ndim == 3 and position_ids.shape[0] == 4:
text_position_ids = position_ids[0]
rotary_position_ids = position_ids[1:]
else:
text_position_ids = None
rotary_position_ids = position_ids
mask_kwargs = {
"config": text_model.config,
"inputs_embeds": inputs_embeds,
"attention_mask": attention_mask,
"past_key_values": None,
"position_ids": text_position_ids,
}
masks = {
"full_attention": create_causal_mask(**mask_kwargs),
"linear_attention": create_recurrent_attention_mask(**mask_kwargs),
}
hidden_states = inputs_embeds
position_embeddings = text_model.rotary_emb(hidden_states, rotary_position_ids)
insertion_index = self.model_config.insertion_after
layers = text_model.layers[: text_model.config.num_hidden_layers]
for layer in layers[: insertion_index + 1]:
hidden_states = layer(
hidden_states,
position_embeddings=position_embeddings,
attention_mask=masks[layer.block_type],
position_ids=text_position_ids,
past_key_values=None,
use_cache=False,
**kwargs,
)
hidden_states = self.reasoning_core(
hidden_states,
position_embeddings=position_embeddings,
masks=masks,
position_ids=text_position_ids,
**kwargs,
)
for layer in layers[insertion_index + 1 :]:
hidden_states = layer(
hidden_states,
position_embeddings=position_embeddings,
attention_mask=masks[layer.block_type],
position_ids=text_position_ids,
past_key_values=None,
use_cache=False,
**kwargs,
)
hidden_states = text_model.norm(hidden_states)
slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep
logits = self.backbone.lm_head(hidden_states[:, slice_indices, :])
loss = None
if labels is not None:
loss = self.backbone.loss_function(
logits=logits,
labels=labels,
vocab_size=self.config.vocab_size,
**kwargs,
)
return CausalLMOutputWithPast(loss=loss, logits=logits, past_key_values=None)
def architecture_manifest(self) -> dict[str, Any]:
return {
"architecture": "DotRecurrentDepthModel",
"base_model": self.model_config.base_model,
"base_parameter_count": sum(parameter.numel() for parameter in self.backbone.parameters()),
"total_parameter_count": self.total_parameter_count(),
"new_parameter_count": sum(
parameter.numel() for parameter in self.reasoning_core.parameters()
),
"model_config": asdict(self.model_config),
}