| """ |
| rLoRA Model: Retrofitted Recurrence with LoRA adapters. |
| |
| Splits a pretrained transformer into prelude/recurrent/coda segments, |
| injects LoRA into the recurrent block, and loops it N times at inference. |
| """ |
|
|
| import torch |
| import torch.nn as nn |
| import torch.nn.functional as F |
| from torch import Tensor |
| from torch.utils.checkpoint import checkpoint as torch_checkpoint |
| from typing import Optional, List, Tuple, Dict, Any, Union |
|
|
| from transformers import ( |
| PreTrainedModel, |
| GenerationMixin, |
| AutoModelForCausalLM, |
| AutoConfig, |
| ) |
| from transformers.modeling_outputs import CausalLMOutputWithPast |
| from transformers.masking_utils import create_causal_mask, create_sliding_window_causal_mask |
|
|
| from .config import RLoRAConfig |
| from .lora_modules import ConcatAdapter |
|
|
|
|
| class RLoRAModel(PreTrainedModel, GenerationMixin): |
| """ |
| Retrofitted Recurrence with LoRA. |
| |
| Takes a pretrained causal LM, splits its layers into prelude/recurrent/coda, |
| injects LoRA into the recurrent block layers, and loops them N times. |
| Only LoRA parameters and the inter-recurrence adapter are trained. |
| """ |
|
|
| config_class = RLoRAConfig |
| base_model_prefix = "rlora" |
| supports_gradient_checkpointing = True |
|
|
| @classmethod |
| def from_pretrained(cls, pretrained_model_name_or_path, *model_args, **kwargs): |
| """Override to avoid meta-device initialization. |
| |
| The default from_pretrained creates the model on meta device then loads |
| weights, but our __init__ calls AutoModelForCausalLM.from_pretrained |
| internally, which fails on meta device. Instead, we build the model |
| normally (loading the base model), then overwrite with saved weights. |
| """ |
| import os |
| from safetensors.torch import load_file |
|
|
| config = kwargs.pop("config", None) |
| if config is None: |
| config, model_kwargs = cls.config_class.from_pretrained( |
| pretrained_model_name_or_path, |
| return_unused_kwargs=True, |
| **kwargs, |
| ) |
| else: |
| model_kwargs = kwargs |
|
|
| |
| model = cls(config) |
|
|
| |
| if os.path.isdir(pretrained_model_name_or_path): |
| weight_path = os.path.join(pretrained_model_name_or_path, "model.safetensors") |
| else: |
| from huggingface_hub import hf_hub_download |
| weight_path = hf_hub_download(pretrained_model_name_or_path, "model.safetensors") |
|
|
| state_dict = load_file(weight_path) |
| model.load_state_dict(state_dict, strict=True) |
|
|
| device = model_kwargs.get("device_map", None) |
| if device is not None: |
| model = model.to(device) |
|
|
| return model |
|
|
| def __init__(self, config: RLoRAConfig): |
| super().__init__(config) |
|
|
| |
| |
| |
| |
| |
| base_config = AutoConfig.from_pretrained( |
| config.base_model_name, trust_remote_code=True |
| ) |
|
|
| base = AutoModelForCausalLM.from_pretrained( |
| config.base_model_name, |
| config=base_config, |
| torch_dtype="auto", |
| trust_remote_code=True, |
| attn_implementation="sdpa", |
| low_cpu_mem_usage=True, |
| ) |
|
|
| for param in base.parameters(): |
| param.requires_grad = False |
| base.eval() |
|
|
| |
| self.base_config = base_config |
| self.hidden_size = base_config.hidden_size |
| self.vocab_size = base_config.vocab_size |
|
|
| |
| self._inject_lora(config, base) |
|
|
| |
| |
| del base |
|
|
| |
| model_dtype = self.embed_tokens.weight.dtype |
| if config.adapter_type == "concat": |
| self.adapter = ConcatAdapter(self.hidden_size, init_type=config.adapter_init).to(dtype=model_dtype) |
| else: |
| self.adapter = None |
|
|
| |
| if config.freeze_lm_head: |
| for param in self.lm_head.parameters(): |
| param.requires_grad = False |
|
|
| |
| self.gradient_checkpointing = False |
|
|
| def _inject_lora(self, config: RLoRAConfig, base_model: nn.Module): |
| """Inject LoRA adapters into recurrent block layers and extract all components.""" |
| from peft import LoraConfig as PeftLoraConfig, get_peft_model |
|
|
| peft_config = PeftLoraConfig( |
| task_type="CAUSAL_LM", |
| r=config.lora_r, |
| lora_alpha=config.lora_alpha, |
| lora_dropout=config.lora_dropout, |
| target_modules=config.lora_target_modules, |
| layers_to_transform=config.recurrent_layers, |
| bias="none", |
| ) |
|
|
| peft_model = get_peft_model(base_model, peft_config) |
|
|
| |
| |
| inner = peft_model.base_model.model.model |
| all_layers = list(inner.layers) |
|
|
| self.prelude_layers = nn.ModuleList([all_layers[i] for i in config.prelude_layers]) |
| self.recurrent_layers = nn.ModuleList([all_layers[i] for i in config.recurrent_layers]) |
| self.coda_layers = nn.ModuleList([all_layers[i] for i in config.coda_layers]) |
| self.embed_tokens = inner.embed_tokens |
| self.final_norm = inner.norm |
| self.rotary_emb = inner.rotary_emb |
| self.lm_head = peft_model.base_model.model.lm_head |
|
|
| def train(self, mode: bool = True): |
| """Override to keep frozen segments in eval mode (no dropout/batchnorm leakage).""" |
| super().train(mode) |
| |
| |
| if hasattr(self, 'prelude_layers'): |
| self.prelude_layers.eval() |
| self.coda_layers.eval() |
| self.embed_tokens.eval() |
| self.final_norm.eval() |
| return self |
|
|
| def print_trainable_parameters(self): |
| """Print trainable vs total parameter counts.""" |
| total = sum(p.numel() for p in self.parameters()) |
| trainable = sum(p.numel() for p in self.parameters() if p.requires_grad) |
| print( |
| f"trainable params: {trainable:,} || all params: {total:,} " |
| f"|| trainable%: {100 * trainable / total:.4f}" |
| ) |
|
|
| def initialize_state(self, input_embeds: Tensor, scale: float = 1.0) -> Tensor: |
| """Initialize recurrent state as truncated normal noise (following the paper).""" |
| x = torch.randn_like(input_embeds) |
| std = self.config.state_init_std * scale |
| if std > 0: |
| torch.nn.init.trunc_normal_(x, mean=0.0, std=std, a=-3 * std, b=3 * std) |
| else: |
| x.zero_() |
| return x |
|
|
| def _build_causal_masks( |
| self, |
| inputs_embeds: Tensor, |
| attention_mask: Optional[Tensor], |
| cache_position: Tensor, |
| past_key_values=None, |
| position_ids: Optional[Tensor] = None, |
| ) -> Dict[str, Optional[Tensor]]: |
| """Build causal mask mapping for OLMo-3's mixed sliding/full attention.""" |
| mask_kwargs = dict( |
| config=self.base_config, |
| inputs_embeds=inputs_embeds, |
| attention_mask=attention_mask, |
| cache_position=cache_position, |
| past_key_values=past_key_values, |
| position_ids=position_ids, |
| ) |
|
|
| has_sliding = any( |
| hasattr(layer, "self_attn") and getattr(layer.self_attn, "attention_type", "full_attention") == "sliding_attention" |
| for layer in list(self.prelude_layers) + list(self.recurrent_layers) + list(self.coda_layers) |
| ) |
|
|
| result = { |
| "full_attention": create_causal_mask(**mask_kwargs), |
| } |
| if has_sliding: |
| result["sliding_attention"] = create_sliding_window_causal_mask(**mask_kwargs) |
|
|
| return result |
|
|
| def _get_attention_type(self, layer: nn.Module) -> str: |
| """Get the attention type for a decoder layer.""" |
| if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "attention_type"): |
| return layer.self_attn.attention_type |
| return "full_attention" |
|
|
| def _run_layer( |
| self, |
| layer: nn.Module, |
| hidden_states: Tensor, |
| causal_mask_mapping: Dict[str, Optional[Tensor]], |
| position_embeddings: Tuple[Tensor, Tensor], |
| position_ids: Optional[Tensor] = None, |
| past_key_values=None, |
| use_cache: bool = False, |
| cache_position: Optional[Tensor] = None, |
| ) -> Tensor: |
| """Run a single decoder layer with the correct attention mask.""" |
| attn_type = self._get_attention_type(layer) |
| mask = causal_mask_mapping.get(attn_type, causal_mask_mapping.get("full_attention")) |
|
|
| if self.gradient_checkpointing and self.training: |
| return torch_checkpoint( |
| layer, |
| hidden_states, |
| attention_mask=mask, |
| position_embeddings=position_embeddings, |
| position_ids=position_ids, |
| past_key_values=past_key_values, |
| use_cache=use_cache, |
| cache_position=cache_position, |
| use_reentrant=False, |
| ) |
|
|
| return layer( |
| hidden_states, |
| attention_mask=mask, |
| position_embeddings=position_embeddings, |
| position_ids=position_ids, |
| past_key_values=past_key_values, |
| use_cache=use_cache, |
| cache_position=cache_position, |
| ) |
|
|
| def _run_layers( |
| self, |
| layers: nn.ModuleList, |
| hidden_states: Tensor, |
| causal_mask_mapping: Dict[str, Optional[Tensor]], |
| position_embeddings: Tuple[Tensor, Tensor], |
| position_ids: Optional[Tensor] = None, |
| past_key_values=None, |
| use_cache: bool = False, |
| cache_position: Optional[Tensor] = None, |
| ) -> Tensor: |
| """Run a sequence of decoder layers.""" |
| for layer in layers: |
| hidden_states = self._run_layer( |
| layer, hidden_states, causal_mask_mapping, |
| position_embeddings, position_ids, |
| past_key_values, use_cache, cache_position, |
| ) |
| return hidden_states |
|
|
| def _single_recurrence_step( |
| self, |
| state: Tensor, |
| prelude_output: Tensor, |
| causal_mask_mapping: Dict[str, Optional[Tensor]], |
| position_embeddings: Tuple[Tensor, Tensor], |
| position_ids: Optional[Tensor] = None, |
| ) -> Tensor: |
| """Execute one recurrence step: inject prelude features + run recurrent block.""" |
| |
| if self.config.adapter_type == "concat": |
| state = self.adapter(state, prelude_output) |
| else: |
| state = state + prelude_output |
|
|
| |
| state = self._run_layers( |
| self.recurrent_layers, state, causal_mask_mapping, |
| position_embeddings, position_ids, |
| past_key_values=None, use_cache=False, |
| ) |
| return state |
|
|
| def forward( |
| self, |
| input_ids: Tensor, |
| attention_mask: Optional[Tensor] = None, |
| position_ids: Optional[Tensor] = None, |
| past_key_values=None, |
| labels: Optional[Tensor] = None, |
| use_cache: bool = False, |
| num_recurrences: Optional[int] = None, |
| num_steps_no_grad: Optional[int] = None, |
| num_steps_with_grad: Optional[int] = None, |
| **kwargs, |
| ) -> Union[Tuple, CausalLMOutputWithPast]: |
| """ |
| Forward pass with recurrent depth and truncated BPTT. |
| |
| Args: |
| num_recurrences: Total recurrence steps (ignored if split provided) |
| num_steps_no_grad: Steps WITHOUT gradients (truncated BPTT) |
| num_steps_with_grad: Steps WITH gradients |
| """ |
| batch_size, seq_len = input_ids.shape |
| device = input_ids.device |
|
|
| |
| if num_steps_no_grad is not None and num_steps_with_grad is not None: |
| total_recurrences = num_steps_no_grad + num_steps_with_grad |
| elif num_recurrences is not None: |
| num_steps_no_grad = 0 |
| num_steps_with_grad = num_recurrences |
| total_recurrences = num_recurrences |
| else: |
| num_steps_no_grad = 0 |
| num_steps_with_grad = self.config.default_num_recurrences |
| total_recurrences = self.config.default_num_recurrences |
|
|
| |
| input_embeds = self.embed_tokens(input_ids) |
|
|
| |
| cache_position = torch.arange(seq_len, device=device) |
| if position_ids is None: |
| position_ids = cache_position.unsqueeze(0).expand(batch_size, -1) |
|
|
| |
| position_embeddings = self.rotary_emb(input_embeds, position_ids) |
|
|
| |
| causal_mask_mapping = self._build_causal_masks( |
| input_embeds, attention_mask, cache_position, |
| past_key_values=None, position_ids=position_ids, |
| ) |
|
|
| |
| hidden_states = self._run_layers( |
| self.prelude_layers, input_embeds, causal_mask_mapping, |
| position_embeddings, position_ids, |
| ) |
| prelude_output = hidden_states |
|
|
| |
| total_recurrences = num_steps_no_grad + num_steps_with_grad |
| if total_recurrences == 1: |
| |
| state = self._run_layers( |
| self.recurrent_layers, prelude_output, causal_mask_mapping, |
| position_embeddings, position_ids, |
| past_key_values=None, use_cache=False, |
| ) |
| else: |
| |
| state = self.initialize_state(prelude_output) |
|
|
| |
| with torch.no_grad(): |
| for _t in range(num_steps_no_grad): |
| state = self._single_recurrence_step( |
| state, prelude_output, causal_mask_mapping, |
| position_embeddings, position_ids, |
| ) |
|
|
| |
| for _t in range(num_steps_with_grad): |
| state = self._single_recurrence_step( |
| state, prelude_output, causal_mask_mapping, |
| position_embeddings, position_ids, |
| ) |
|
|
| |
| hidden_states = self._run_layers( |
| self.coda_layers, state, causal_mask_mapping, |
| position_embeddings, position_ids, |
| ) |
|
|
| hidden_states = self.final_norm(hidden_states) |
|
|
| |
| logits = self.lm_head(hidden_states) |
|
|
| |
| loss = None |
| if labels is not None: |
| |
| |
| |
| shift_logits = logits[..., :-1, :].reshape(-1, self.vocab_size) |
| shift_labels = labels[..., 1:].reshape(-1) |
| loss = F.cross_entropy( |
| shift_logits.float(), shift_labels, ignore_index=-100, |
| ) |
|
|
| return CausalLMOutputWithPast( |
| loss=loss, |
| logits=None if (self.training and labels is not None) else logits, |
| past_key_values=None, |
| hidden_states=None, |
| attentions=None, |
| ) |
|
|
| def prepare_inputs_for_generation( |
| self, |
| input_ids: Tensor, |
| past_key_values=None, |
| attention_mask: Optional[Tensor] = None, |
| inputs_embeds: Optional[Tensor] = None, |
| **kwargs, |
| ) -> Dict[str, Any]: |
| """Prepare inputs for model.generate().""" |
| |
| position_ids = kwargs.get("position_ids", None) |
| if position_ids is None and attention_mask is not None: |
| position_ids = attention_mask.long().cumsum(-1) - 1 |
| position_ids.masked_fill_(attention_mask == 0, 1) |
|
|
| return { |
| "input_ids": input_ids, |
| "attention_mask": attention_mask, |
| "position_ids": position_ids, |
| "past_key_values": None, |
| "use_cache": False, |
| "num_recurrences": kwargs.get("num_recurrences", self.config.default_num_recurrences), |
| } |
|
|