talkie-1930-13b-it-gptq-int4 / modeling_talkie.py
dtestnyrr's picture
Initial GPTQ int4 upload
79636b3 verified
Raw History Blame Contribute Delete
10.2 kB
"""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)