rlora-olmo3-7b / lora_modules.py
jhdlee's picture
Upload folder using huggingface_hub
8a84a9b verified
Raw
History Blame Contribute Delete
1.34 kB
"""Learnable modules for rLoRA inter-recurrence adaptation."""
import torch
import torch.nn as nn
class ConcatAdapter(nn.Module):
"""
Inter-recurrence adapter: Linear(cat([state, prelude_output], dim=-1)).
Two initialization modes:
- "default": PyTorch Kaiming uniform with bias (matches the paper).
- "prelude_identity": Identity on prelude half, zeros on state half,
no bias. Output ≈ prelude_output at init, giving pretrained-compatible
input to the recurrent block from the start.
"""
def __init__(self, hidden_size: int, init_type: str = "default"):
super().__init__()
self.linear = nn.Linear(hidden_size * 2, hidden_size, bias=(init_type == "default"))
if init_type == "prelude_identity":
self._prelude_identity_init(hidden_size)
# else: keep PyTorch default Kaiming init (matches paper)
def _prelude_identity_init(self, hidden_size: int):
"""Initialize so output ≈ second half of input (the prelude output)."""
with torch.no_grad():
self.linear.weight.zero_()
self.linear.weight[:, hidden_size:].copy_(torch.eye(hidden_size))
def forward(self, state: torch.Tensor, prelude_output: torch.Tensor) -> torch.Tensor:
return self.linear(torch.cat([state, prelude_output], dim=-1))