"""HuggingFace-compatible port of TalkieModel. Mirrors talkie/src/talkie/model.py exactly: - F.rms_norm everywhere (no learnable RMSNorm scale) - RoPE with base=1e6 - Attention with QK-norm and per-head HeadGain on Q - SwiGLU MLP - Per-layer ActGain on attn / mlp / embed_skip residuals - WeightGain on the lm_head matrix Linear projections are renamed to Llama conventions (q_proj, k_proj, v_proj, o_proj, gate_proj, up_proj, down_proj) so GPTQModel auto-detects them. """ from __future__ import annotations import torch import torch.nn as nn import torch.nn.functional as F from transformers import PreTrainedModel from transformers.generation import GenerationMixin from transformers.modeling_outputs import CausalLMOutput from .configuration_talkie import TalkieConfig def _apply_rotary_emb(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: # x: [B, T, H, D]; cos/sin: [1, T, 1, D/2] d = x.shape[-1] // 2 x1, x2 = x[..., :d], x[..., d:] y1 = x1 * cos + x2 * sin y2 = -x1 * sin + x2 * cos return torch.cat([y1, y2], dim=-1).type_as(x) class HeadGain(nn.Module): def __init__(self, n_head: int): super().__init__() self.head_g = nn.Parameter(torch.ones(n_head)) def forward(self, x: torch.Tensor) -> torch.Tensor: return x * self.head_g.type_as(x).view(1, 1, -1, 1) class WeightGain(nn.Module): def __init__(self): super().__init__() self.w_g = nn.Parameter(torch.ones(1)) def forward(self, w: torch.Tensor) -> torch.Tensor: return w * self.w_g.type_as(w) class ActGain(nn.Module): def __init__(self, init_value: float): super().__init__() self.a_g = nn.Parameter(torch.ones(1) * init_value) def forward(self, x: torch.Tensor) -> torch.Tensor: return x * self.a_g.type_as(x) class TalkieAttention(nn.Module): def __init__(self, config: TalkieConfig): super().__init__() self.n_head = config.num_attention_heads self.head_dim = config.head_dim h = config.hidden_size self.q_proj = nn.Linear(h, h, bias=False) self.k_proj = nn.Linear(h, h, bias=False) self.v_proj = nn.Linear(h, h, bias=False) self.o_proj = nn.Linear(h, h, bias=False) self.head_gain = HeadGain(config.num_attention_heads) def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: bsz, seq_len, _ = x.size() q = self.q_proj(x).view(bsz, seq_len, self.n_head, self.head_dim) k = self.k_proj(x).view(bsz, seq_len, self.n_head, self.head_dim) v = self.v_proj(x).view(bsz, seq_len, self.n_head, self.head_dim) q = _apply_rotary_emb(q, cos, sin) k = _apply_rotary_emb(k, cos, sin) q = F.rms_norm(q, (q.size(-1),)) k = F.rms_norm(k, (k.size(-1),)) q = self.head_gain(q) # SDPA expects [B, H, T, D] y = F.scaled_dot_product_attention( q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), is_causal=True ) y = y.transpose(1, 2).contiguous().view(bsz, seq_len, -1) return self.o_proj(y) class TalkieMLP(nn.Module): def __init__(self, config: TalkieConfig): super().__init__() self.gate_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) self.up_proj = nn.Linear(config.hidden_size, config.intermediate_size, bias=False) self.down_proj = nn.Linear(config.intermediate_size, config.hidden_size, bias=False) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) class TalkieDecoderLayer(nn.Module): def __init__(self, config: TalkieConfig): super().__init__() self.self_attn = TalkieAttention(config) self.attn_gain = ActGain((2 * config.num_hidden_layers) ** -0.5) self.mlp = TalkieMLP(config) self.mlp_gain = ActGain((2 * config.num_hidden_layers) ** -0.5) self.embed_skip = ActGain(0.0) def forward( self, hidden_states: torch.Tensor, e_x: torch.Tensor = None, cos: torch.Tensor = None, sin: torch.Tensor = None, **kwargs, # absorb attention_mask / position_ids / etc that HF tooling injects ) -> torch.Tensor: x = hidden_states x = x + self.attn_gain(self.self_attn(F.rms_norm(x, (x.shape[-1],)), cos, sin)) x = x + self.mlp_gain(self.mlp(F.rms_norm(x, (x.shape[-1],)))) x = x + self.embed_skip(e_x) return x class TalkiePreTrainedModel(PreTrainedModel): config_class = TalkieConfig base_model_prefix = "model" supports_gradient_checkpointing = False _no_split_modules = ["TalkieDecoderLayer"] def _init_weights(self, module): if isinstance(module, nn.Linear): module.weight.data.normal_(mean=0.0, std=0.02) if module.bias is not None: module.bias.data.zero_() elif isinstance(module, nn.Embedding): module.weight.data.normal_(mean=0.0, std=0.02) class TalkieModel(TalkiePreTrainedModel): def __init__(self, config: TalkieConfig): super().__init__(config) self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size) self.layers = nn.ModuleList( [TalkieDecoderLayer(config) for _ in range(config.num_hidden_layers)] ) # cos/sin are computed lazily in forward — see _rope. Avoid register_buffer # so HF's meta-init / low_cpu_mem_usage loading path does not leave us # holding meta tensors that we then try to slice (which raises). self._rope_cache: tuple[torch.Tensor, torch.Tensor, torch.device, torch.dtype, int] | None = None self.post_init() @staticmethod def _build_rope( seq_len: int, head_dim: int, base: float, device, dtype ) -> tuple[torch.Tensor, torch.Tensor]: ch = torch.arange(0, head_dim, 2, dtype=torch.float32, device=device) inv_freq = 1.0 / (base ** (ch / head_dim)) t = torch.arange(seq_len, dtype=torch.float32, device=device) freqs = torch.outer(t, inv_freq) cos, sin = freqs.cos().to(dtype), freqs.sin().to(dtype) return cos[None, :, None, :], sin[None, :, None, :] def _rope(self, seq_len: int, device, dtype) -> tuple[torch.Tensor, torch.Tensor]: cache = self._rope_cache if (cache is None or cache[2] != device or cache[3] != dtype or cache[4] < seq_len): cap = max(seq_len, self.config.max_position_embeddings) cos, sin = self._build_rope(cap, self.config.head_dim, self.config.rope_theta, device, dtype) self._rope_cache = (cos, sin, device, dtype, cap) cos, sin, _, _, _ = self._rope_cache return cos[:, :seq_len], sin[:, :seq_len] def forward(self, input_ids: torch.LongTensor, **kwargs) -> torch.Tensor: _, seq_len = input_ids.shape x = self.embed_tokens(input_ids) x = F.rms_norm(x, (x.shape[-1],)) e_x = x # post-RMSNorm input embeddings; reused as embed_skip source at every layer cos, sin = self._rope(seq_len, x.device, x.dtype) for layer in self.layers: # Pass e_x/cos/sin as kwargs so HF tooling (GPTQModel etc) captures # and replays them per-sample when iterating layers individually. x = layer(x, e_x=e_x, cos=cos, sin=sin) x = F.rms_norm(x, (x.shape[-1],)) return x class TalkieForCausalLM(TalkiePreTrainedModel, GenerationMixin): _tied_weights_keys = [] _supports_cache_class = False _supports_static_cache = False # Talkie has no KV cache implementation — every generate step recomputes # the full sequence. Mirror the reference talkie inference behavior. def __init__(self, config: TalkieConfig): super().__init__(config) # Force use_cache=False so HF generate doesn't try to feed only the # last token via past_key_values (which we don't support). config.use_cache = False self.model = TalkieModel(config) self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.lm_head_gain = WeightGain() self.post_init() # Belt-and-suspenders: also force the generation_config to not cache. if hasattr(self, "generation_config") and self.generation_config is not None: self.generation_config.use_cache = False def prepare_inputs_for_generation(self, input_ids, **kwargs): # Strip past_key_values and always feed the full sequence — talkie # has no incremental state. kwargs.pop("past_key_values", None) kwargs.pop("cache_position", None) kwargs["use_cache"] = False return {"input_ids": input_ids, **kwargs} 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_embeddings): self.lm_head = new_embeddings def forward( self, input_ids: torch.LongTensor = None, attention_mask: torch.Tensor = None, # accepted but unused (causal-only) labels: torch.LongTensor = None, **kwargs, ) -> CausalLMOutput: hidden = self.model(input_ids) # WeightGain is scalar-broadcast over the lm_head matrix, so applying # it on the linear's output is mathematically identical to pre-scaling # the weight (and avoids a cross-module tensor passing pattern that # confuses accelerate's device-map hooks). logits = self.lm_head(hidden).float() * self.lm_head_gain.w_g.float() 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), ) return CausalLMOutput(loss=loss, logits=logits)