babyloop-loop-text / modeling_babyloop.py
daichi812's picture
final model (main)
6a0cf5a verified
Raw
History Blame Contribute Delete
10.8 kB
"""ループドTransformer本体(HF ``PreTrainedModel`` 互換・単一実装)。
設計 docs/architecture.md §2 の Llama 系レシピ(RMSNorm / RoPE / SwiGLU /
bias なし / weight tying)を素の PyTorch で実装。``forward`` は K 回ループする
形で書き、``k=1`` で標準Transformerに厳密に縮退する(``tests/test_k1_equivalence.py``)。
このファイルは ``save_pretrained`` 時に checkpoint へ複製され、公式
evaluation-pipeline(別プロセス・``trust_remote_code=True``)から import される。
そのため **torch / transformers / 標準ライブラリ以外に依存しない**こと。
probing 用のループ毎中間表現は HF 標準 ``hidden_states``(層ごと)と混ぜず、
別フィールド ``loop_hidden_states`` に格納する。
"""
from __future__ import annotations
import math
from dataclasses import dataclass
import torch
import torch.nn.functional as F
from torch import nn
from transformers.modeling_outputs import ModelOutput
from transformers.modeling_utils import PreTrainedModel
from .configuration_babyloop import BabyloopConfig
# --- ビルディングブロック(Llama系レシピ)---------------------------------
class RMSNorm(nn.Module):
def __init__(self, dim: int, eps: float):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x: torch.Tensor) -> torch.Tensor:
dtype = x.dtype
x = x.float()
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
return (x.to(dtype)) * self.weight
def _rotate_half(x: torch.Tensor) -> torch.Tensor:
x1, x2 = x.chunk(2, dim=-1)
return torch.cat((-x2, x1), dim=-1)
def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor:
# x: (B, H, T, head_dim); cos/sin: (1, 1, T, head_dim)
return x * cos + _rotate_half(x) * sin
class Attention(nn.Module):
"""RoPE 付き causal Multi-Head Attention(dropout なし)。"""
def __init__(self, config: BabyloopConfig):
super().__init__()
self.n_heads = config.n_heads
self.head_dim = config.d_model // config.n_heads
self.qkv_proj = nn.Linear(config.d_model, 3 * config.d_model, bias=config.bias)
self.o_proj = nn.Linear(config.d_model, config.d_model, bias=config.bias)
def forward(self, x, cos, sin, attn_bias):
B, T, C = x.shape
q, k, v = self.qkv_proj(x).split(C, dim=-1)
q = q.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
k = k.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
v = v.view(B, T, self.n_heads, self.head_dim).transpose(1, 2)
q = _apply_rope(q, cos, sin)
k = _apply_rope(k, cos, sin)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_bias, is_causal=attn_bias is None
)
out = out.transpose(1, 2).reshape(B, T, C)
return self.o_proj(out)
class SwiGLU(nn.Module):
def __init__(self, config: BabyloopConfig):
super().__init__()
self.gate_proj = nn.Linear(config.d_model, config.ffn_hidden, bias=config.bias)
self.up_proj = nn.Linear(config.d_model, config.ffn_hidden, bias=config.bias)
self.down_proj = nn.Linear(config.ffn_hidden, config.d_model, bias=config.bias)
def forward(self, x):
return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
class Block(nn.Module):
"""Pre-Norm 残差ブロック: h += Attn(RMSNorm(h)); h += SwiGLU(RMSNorm(h))。"""
def __init__(self, config: BabyloopConfig):
super().__init__()
self.attn_norm = RMSNorm(config.d_model, config.rms_eps)
self.attn = Attention(config)
self.mlp_norm = RMSNorm(config.d_model, config.rms_eps)
self.mlp = SwiGLU(config)
def forward(self, h, cos, sin, attn_bias):
h = h + self.attn(self.attn_norm(h), cos, sin, attn_bias)
h = h + self.mlp(self.mlp_norm(h))
return h
# --- 出力コンテナ -----------------------------------------------------------
@dataclass
class LoopedModelOutput(ModelOutput):
last_hidden_state: torch.FloatTensor | None = None
hidden_states: tuple[torch.FloatTensor, ...] | None = None
loop_hidden_states: tuple[torch.FloatTensor, ...] | None = None
@dataclass
class LoopedCausalLMOutput(ModelOutput):
loss: torch.FloatTensor | None = None
logits: torch.FloatTensor | None = None
hidden_states: tuple[torch.FloatTensor, ...] | None = None
loop_hidden_states: tuple[torch.FloatTensor, ...] | None = None
# --- PreTrainedModel ラッパ -------------------------------------------------
class LoopedPreTrainedModel(PreTrainedModel):
config_class = BabyloopConfig
base_model_prefix = "model"
supports_gradient_checkpointing = False
def _init_weights(self, module):
std = 0.02
if isinstance(module, nn.Linear):
nn.init.normal_(module.weight, mean=0.0, std=std)
if module.bias is not None:
nn.init.zeros_(module.bias)
elif isinstance(module, nn.Embedding):
nn.init.normal_(module.weight, mean=0.0, std=std)
elif isinstance(module, RMSNorm):
nn.init.ones_(module.weight)
class LoopedModel(LoopedPreTrainedModel):
"""重み共有 core ブロックを K 回反復するバックボーン(lm_head なし)。"""
def __init__(self, config: BabyloopConfig):
super().__init__(config)
self.config = config
self.embed_tokens = nn.Embedding(config.vocab_size, config.d_model)
self.prelude = nn.ModuleList(Block(config) for _ in range(config.n_prelude))
self.core = nn.ModuleList(Block(config) for _ in range(config.n_core))
self.coda = nn.ModuleList(Block(config) for _ in range(config.n_coda))
self.final_norm = RMSNorm(config.d_model, config.rms_eps)
head_dim = config.d_model // config.n_heads
inv_freq = 1.0 / (
config.rope_base ** (torch.arange(0, head_dim, 2, dtype=torch.float32) / head_dim)
)
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.post_init()
def get_input_embeddings(self):
return self.embed_tokens
def set_input_embeddings(self, value):
self.embed_tokens = value
def _rope(self, T: int, device, dtype):
t = torch.arange(T, device=device, dtype=torch.float32)
freqs = torch.outer(t, self.inv_freq.to(device))
emb = torch.cat((freqs, freqs), dim=-1)
return emb.cos().to(dtype)[None, None], emb.sin().to(dtype)[None, None]
def _attn_bias(self, attention_mask, T, device, dtype):
# padding が無ければ None を返し、SDPA の is_causal 経路(高速)に乗せる。
if attention_mask is None or bool((attention_mask == 1).all()):
return None
causal = torch.ones(T, T, device=device, dtype=torch.bool).triu(1)
key_pad = attention_mask.to(device) == 0 # (B, T)
mask = causal[None, None] | key_pad[:, None, None, :]
bias = torch.zeros(mask.shape, device=device, dtype=dtype)
return bias.masked_fill(mask, torch.finfo(dtype).min)
def forward(
self,
input_ids=None,
attention_mask=None,
inputs_embeds=None,
output_hidden_states=False,
**kwargs,
) -> LoopedModelOutput:
if inputs_embeds is None:
inputs_embeds = self.embed_tokens(input_ids)
h = inputs_embeds
residual_input = inputs_embeds # inject_input 用(③で本格化、①は false)
B, T, _ = h.shape
cos, sin = self._rope(T, h.device, h.dtype)
attn_bias = self._attn_bias(attention_mask, T, h.device, h.dtype)
all_hidden = [h] if output_hidden_states else None
loop_hidden = []
for block in self.prelude:
h = block(h, cos, sin, attn_bias)
if output_hidden_states:
all_hidden.append(h)
for _ in range(self.config.k):
for block in self.core:
h = block(h, cos, sin, attn_bias)
if output_hidden_states:
all_hidden.append(h)
if self.config.inject_input:
h = h + residual_input
loop_hidden.append(h)
for block in self.coda:
h = block(h, cos, sin, attn_bias)
if output_hidden_states:
all_hidden.append(h)
h = self.final_norm(h)
return LoopedModelOutput(
last_hidden_state=h,
hidden_states=tuple(all_hidden) if output_hidden_states else None,
loop_hidden_states=tuple(loop_hidden),
)
class LoopedForCausalLM(LoopedPreTrainedModel):
"""言語モデリングヘッド付き(``AutoModelForCausalLM`` 互換)。"""
_tied_weights_keys = ["lm_head.weight"]
def __init__(self, config: BabyloopConfig):
super().__init__(config)
self.model = LoopedModel(config)
self.lm_head = nn.Linear(config.d_model, config.vocab_size, bias=False)
self.post_init()
def get_input_embeddings(self):
return self.model.embed_tokens
def set_input_embeddings(self, value):
self.model.embed_tokens = value
def get_output_embeddings(self):
return self.lm_head
def set_output_embeddings(self, new):
self.lm_head = new
def forward(
self,
input_ids=None,
attention_mask=None,
inputs_embeds=None,
labels=None,
output_hidden_states=False,
**kwargs,
) -> LoopedCausalLMOutput:
out = self.model(
input_ids=input_ids,
attention_mask=attention_mask,
inputs_embeds=inputs_embeds,
output_hidden_states=output_hidden_states,
)
logits = self.lm_head(out.last_hidden_state)
loss = None
if labels is not None:
shift_logits = logits[:, :-1, :].contiguous()
shift_labels = labels[:, 1:].contiguous()
loss = F.cross_entropy(
shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1),
ignore_index=-100,
)
return LoopedCausalLMOutput(
loss=loss,
logits=logits,
hidden_states=out.hidden_states,
loop_hidden_states=out.loop_hidden_states,
)
# 公式 evaluation-pipeline が trust_remote_code で AutoModel 系から読めるよう登録。
# save_pretrained 時に auto_map と本ファイル群が checkpoint へ複製される。
BabyloopConfig.register_for_auto_class()
LoopedModel.register_for_auto_class("AutoModel")
LoopedForCausalLM.register_for_auto_class("AutoModelForCausalLM")