""" FWKV-ROSA: a from-scratch RWKV-style recurrent LM (~56M params) with ROSA copy signal. Self-contained model definition for Hugging Face `transformers`. """ import math from typing import Optional, Tuple, List import torch import torch.nn as nn import torch.nn.functional as F from transformers import ( PretrainedConfig, PreTrainedModel, GenerationMixin, ) from transformers.modeling_outputs import CausalLMOutputWithPast # ------------------------------------------------------------------------- # 1. ROSA (Rapid Online Suffix Automaton) # ------------------------------------------------------------------------- def rosa(x: List[int]) -> List[int]: """ y[i] = x[j+1] where x[j-m:j] == x[i-m:i] for the largest match length m, over all j < i; ties broken by largest j. Returns -1 if no prior occurrence. """ n = len(x) if n == 0: return [] y = [-1] * n s = 2 * n + 2 trans = [None] * s link = [-1] * s length = [0] * s last_end = [-1] * s trans[0] = {} last = 0 size = 1 for i, t in enumerate(x): cur = size size += 1 trans[cur] = {} length[cur] = length[last] + 1 p = last while p != -1 and t not in trans[p]: trans[p][t] = cur p = link[p] if p == -1: link[cur] = 0 else: q = trans[p][t] if length[p] + 1 == length[q]: link[cur] = q else: clone = size size += 1 trans[clone] = trans[q].copy() length[clone] = length[p] + 1 link[clone] = link[q] last_end[clone] = last_end[q] while p != -1 and trans[p][t] == q: trans[p][t] = clone p = link[p] link[q] = clone link[cur] = clone last = cur v = cur pred = -1 while v != -1: if length[v] > 0 and last_end[v] >= 0: pred = x[last_end[v] + 1] break v = link[v] y[i] = pred v = last while v != -1 and last_end[v] < i: last_end[v] = i v = link[v] return y # ------------------------------------------------------------------------- # 2. Parallel scan (exact, vectorised) # ------------------------------------------------------------------------- def parallel_scan_decay(a: torch.Tensor, W: torch.Tensor) -> torch.Tensor: """Hillis-Steele inclusive scan with constant per-channel decay W.""" W = W.to(dtype=a.dtype) val = a T = a.shape[1] d = 1 while d < T: shifted = F.pad(val[:, :-d, :], (0, 0, d, 0)) val = val + (W ** d) * shifted d *= 2 return val # ------------------------------------------------------------------------- # 3. Config # ------------------------------------------------------------------------- class FWKVConfig(PretrainedConfig): model_type = "fwkv" def __init__( self, d_model: int = 512, d_emb: int = 128, n_layers: int = 14, ffn_mult: int = 4, vocab_size: int = 50257, seq_len: int = 1024, wkv_floor: float = 0.1, tie_word_embeddings: bool = True, **kwargs, ): super().__init__(tie_word_embeddings=tie_word_embeddings, **kwargs) self.d_model = d_model self.d_emb = d_emb self.n_layers = n_layers self.ffn_mult = ffn_mult self.vocab_size = vocab_size self.seq_len = seq_len self.wkv_floor = wkv_floor # ------------------------------------------------------------------------- # 4. Factorised tied embedding/head # ------------------------------------------------------------------------- class FactorizedTiedHead(nn.Module): def __init__(self, vocab_size: int, d_model: int, d_emb: int): super().__init__() self.d_model = d_model self.d_emb = d_emb self.weight = nn.Parameter(torch.empty(vocab_size, d_emb)) self.proj = nn.Linear(d_emb, d_model, bias=False) def embed(self, input_ids): return self.proj(F.embedding(input_ids, self.weight)) def to_emb_space(self, x): return F.linear(x, self.proj.weight.t()) def logits(self, x_emb): return F.linear(x_emb, self.weight) # ------------------------------------------------------------------------- # 5. FWKV Block # ------------------------------------------------------------------------- class FWKVBlock(nn.Module): def __init__(self, d: int, ffn_mult: int = 4, floor: float = 0.1): super().__init__() self.floor = floor self.proj_k = nn.Linear(d, d, bias=False) self.proj_v = nn.Linear(d, d, bias=False) self.proj_r = nn.Linear(d, d, bias=False) self.proj_out = nn.Linear(d, d, bias=False) self.w = nn.Parameter(torch.ones(d) * 2.0) self.ffn = nn.Sequential( nn.Linear(d, ffn_mult * d, bias=False), nn.GELU(), nn.Linear(ffn_mult * d, d, bias=False), ) self.norm_wkv = nn.LayerNorm(d) self.norm_ffn = nn.LayerNorm(d) @property def W(self): return torch.clamp(torch.sigmoid(self.w), min=self.floor) def forward(self, x, state=None): B, T, d = x.shape W = self.W k = self.proj_k(x) v = self.proj_v(x) r = torch.sigmoid(self.proj_r(x)) a = k * v if state is not None: a = a.clone() a[:, 0] = a[:, 0] + W * state wkv_out = parallel_scan_decay(a, W) new_state = wkv_out[:, -1].detach() x = self.norm_wkv(x + self.proj_out(r * wkv_out)) x = self.norm_ffn(x + self.ffn(x)) return x, new_state # ------------------------------------------------------------------------- # 6. Full Language Model # ------------------------------------------------------------------------- class FWKVLanguageModel(PreTrainedModel, GenerationMixin): config_class = FWKVConfig base_model_prefix = "fwkv" supports_gradient_checkpointing = True def __init__(self, config: FWKVConfig): super().__init__(config) self.shared = FactorizedTiedHead(config.vocab_size, config.d_model, config.d_emb) self.rosa_emb = nn.Embedding(config.vocab_size + 1, config.d_emb, padding_idx=0) self.blocks = nn.ModuleList([ FWKVBlock(config.d_model, config.ffn_mult, config.wkv_floor) for _ in range(config.n_layers) ]) self.norm = nn.LayerNorm(config.d_model) self.post_init() def get_input_embeddings(self): return self.shared.weight def set_input_embeddings(self, value): self.shared.weight = nn.Parameter(value) def forward( self, input_ids: torch.LongTensor, rosa_ids: Optional[torch.LongTensor] = None, past_key_values: Optional[List[Optional[torch.Tensor]]] = None, labels: Optional[torch.LongTensor] = None, use_cache: bool = True, **kwargs, ) -> CausalLMOutputWithPast: # ROSA IDs: compute if not supplied if rosa_ids is None: rosa_ids = torch.tensor( [rosa(row.tolist()) for row in input_ids.detach().cpu()], device=input_ids.device, dtype=torch.long ) x = self.shared.embed(input_ids) # [B, T, d_model] rosa_idx = (rosa_ids + 1).clamp(min=0) # -1 -> 0 x = x + self.shared.proj(self.rosa_emb(rosa_idx)) # inject ROSA signal # State propagation states_in = past_key_values or [None] * len(self.blocks) states_out = [] for block, state in zip(self.blocks, states_in): x, new_state = block(x, state) states_out.append(new_state) x = self.norm(x) x_emb = self.shared.to_emb_space(x) # [B, T, d_emb] logits = self.shared.logits(x_emb) # [B, T, vocab] loss = None if labels is not None: # Shift logits and labels for standard LM loss 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 CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=states_out if use_cache else None, ) def prepare_inputs_for_generation( self, input_ids: torch.LongTensor, past_key_values: Optional[List[Optional[torch.Tensor]]] = None, rosa_ids: Optional[torch.LongTensor] = None, **kwargs, ): if past_key_values is not None: # Only keep last token for autoregressive step input_ids = input_ids[:, -1:] if rosa_ids is not None: rosa_ids = rosa_ids[:, -1:] return { "input_ids": input_ids, "rosa_ids": rosa_ids, "past_key_values": past_key_values, "use_cache": True, } def _reorder_cache(self, past_key_values, beam_idx): """Reorder past states if beam search is used (unlikely).""" reordered_past = [] for state in past_key_values: if state is not None: reordered_past.append(state.index_select(0, beam_idx)) else: reordered_past.append(None) return reordered_past