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), }