"""TWLAT-R(V2)模型(《04 實驗設計》§4)。 任務:對文本中每個 proposal(字典判定「這裡可能要改」)在候選集中選一個。 候選 index 0 **永遠是「維持原樣」**,即文本中實際出現的形式。 **核心約束(V2 的全部重點)**: 禁止任何 per-candidate / per-site 的 trainable embedding lookup。 候選只能由「表面字串 + 數值特徵」動態編碼,且與文本共用同一份 char embedding。 因此新增字典條目不需要新增任何參數,held-out proposal 也不是隨機向量。 (V1 的 `cand_emb = nn.Embedding(4096, 192)` 佔 13% 參數並阻斷 zero-shot,就是要修掉的。) 結構: char_emb(共用)──┬─→ context encoder ──→ span mean-pool ──→ h_i [B,P,D] │ └─→ 候選字串 mean-pool ──→ proj ──────────→ e_c [B,P,C,D] score = MLP([h_i ⊕ e_c ⊕ (h_i * e_c) ⊕ cand_feat]) → [B,P,C] context encoder 可切換(`TWLATRConfig.encoder`),兩者參數量刻意對齊以便做 scaling curve: - "tcn" :4 層 dilated depthwise separable Conv1d(dilation 1/2/4/8、kernel 5) - "local" :4 層 local-window Transformer(window 半徑 32、RoPE、pre-LN、GELU) batch 欄位見 `TWLATR.forward` docstring。 """ from __future__ import annotations from dataclasses import dataclass from typing import Any import torch import torch.nn as nn import torch.nn.functional as F # Loss 權重(《04》§0 的錯誤分析:73% 的錯是「改了不該改」,故 keep_bias 直接壓它) W_CAND = 1.0 W_KEEP_BIAS = 0.3 KEEP_MARGIN = 0.5 KEEP_INDEX = 0 # 候選 0 恆為「維持原樣」 ENCODERS = ("tcn", "local") @dataclass class TWLATRConfig: """TWLAT-R 配置。目標參數量 3.0M–4.5M。 註:d_model=160 只有約 2.5M(低於下限),故預設起跳為 192; tools/param_count_r.py 會印出兩者的對照。 """ d_model: int = 192 n_heads: int = 4 ffn_dim: int = 768 dropout: float = 0.1 # context encoder:可切換,兩種都必須可跑 encoder: str = "local" n_layers: int = 4 local_window: int = 32 # 半徑;|i-j| <= 32 才可見 conv_kernel: int = 5 conv_dilations: tuple[int, ...] = (1, 2, 4, 8) # depthwise 之後的 pointwise 通道擴張倍率;設 1 即教科書式 depthwise separable, # 設 2 可讓 tcn 與 local 的 mixer 參數量幾乎相等(scaling curve 才公平) conv_expansion: int = 2 # char embedding(文本與候選共用;無任何 per-ID 候選表) han_vocab: int = 4000 # 常用字直接查表;id >= han_vocab 走 hash n_hash: int = 2 hash_buckets: int = 2048 feat_dim: int = 4 # script / 是否在 proposal span 內 / 是否保護段 / 詞界 feat_vocab: tuple[int, ...] = (16, 4, 4, 4) # proposal / candidate 形狀 max_props: int = 48 # P max_cands: int = 8 # C cand_len: int = 6 # L,候選表面字串的最大字元數 cand_feat_dim: int = 12 # K score_hidden: int = 512 # 候選表徵方式: # "dynamic" —— 由表面字串經共用 char_emb 組合而成(V2 預設,可 zero-shot) # "id" —— per-candidate learned lookup(H-B 的對照組,複製 V1 的失敗模式) candidate_encoder: str = "dynamic" cand_id_vocab: int = 4096 seq_len: int = 512 rope_base: float = 10000.0 def __post_init__(self) -> None: assert self.encoder in ENCODERS, f"encoder 必須是 {ENCODERS}" assert self.d_model % self.n_heads == 0 assert self.d_model % 2 == 0, "hash_dim = d_model // 2,需為偶數" assert len(self.feat_vocab) == self.feat_dim assert self.conv_kernel % 2 == 1, "kernel 需為奇數才能等長 padding" @property def head_dim(self) -> int: return self.d_model // self.n_heads @property def hash_dim(self) -> int: return self.d_model // 2 def dilation_at(self, i: int) -> int: return self.conv_dilations[i % len(self.conv_dilations)] # --------------------------------------------------------------------------- # # RoPE # --------------------------------------------------------------------------- # def build_rope_cache( seq_len: int, head_dim: int, base: float, device, dtype ) -> tuple[torch.Tensor, torch.Tensor]: """回傳 [T, head_dim//2] 的 cos / sin。""" half = head_dim // 2 inv_freq = base ** (-torch.arange(half, device=device, dtype=torch.float32) / half) pos = torch.arange(seq_len, device=device, dtype=torch.float32) freqs = torch.outer(pos, inv_freq) return freqs.cos().to(dtype), freqs.sin().to(dtype) def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: """x: [B, H, T, D],對相鄰兩維做旋轉。""" x_even, x_odd = x[..., 0::2], x[..., 1::2] cos = cos[None, None, : x.shape[-2], :] sin = sin[None, None, : x.shape[-2], :] out = torch.stack([x_even * cos - x_odd * sin, x_even * sin + x_odd * cos], dim=-1) return out.flatten(-2) # --------------------------------------------------------------------------- # # Embedding:文本與候選共用 # --------------------------------------------------------------------------- # _HASH_MULT = (2654435761, 40503) _HASH_ADD = (0, 987654321) class CharEmbedding(nn.Module): """常用字 4000 直接查表;id >= han_vocab 的罕字用 2 組 hash(buckets 2048, dim d/2) 串接後投影。 **文本與候選字串共用這一份**:候選只是一串字元 id,沒有自己的 embedding 表, 所以字典新增條目不會增加任何參數,未見過的字串也落在同一個表徵空間。 """ def __init__(self, cfg: TWLATRConfig): super().__init__() self.cfg = cfg self.han = nn.Embedding(cfg.han_vocab, cfg.d_model) self.hash = nn.ModuleList( nn.Embedding(cfg.hash_buckets, cfg.hash_dim) for _ in range(cfg.n_hash) ) self.hash_proj = nn.Linear(cfg.n_hash * cfg.hash_dim, cfg.d_model) def forward(self, ids: torch.Tensor) -> torch.Tensor: """ids: 任意形狀 [...],回傳 [..., d_model]。""" cfg = self.cfg ids = ids.clamp(min=0) rare = ids >= cfg.han_vocab han = self.han(ids.clamp(max=cfg.han_vocab - 1)) parts = [] for i, emb in enumerate(self.hash): m = _HASH_MULT[i % len(_HASH_MULT)] a = _HASH_ADD[i % len(_HASH_ADD)] parts.append(emb((ids * m + a) % cfg.hash_buckets)) rare_vec = self.hash_proj(torch.cat(parts, dim=-1)) return torch.where(rare.unsqueeze(-1), rare_vec, han) class TextFeatures(nn.Module): """文本側的 4 個離散特徵;char embedding 由外部傳入,以免共用的表被重複註冊。""" def __init__(self, cfg: TWLATRConfig): super().__init__() self.cfg = cfg self.feat = nn.ModuleList(nn.Embedding(n, cfg.d_model) for n in cfg.feat_vocab) self.ln = nn.LayerNorm(cfg.d_model) self.drop = nn.Dropout(cfg.dropout) def forward(self, x: torch.Tensor, feat: torch.Tensor) -> torch.Tensor: for i, emb in enumerate(self.feat): x = x + emb(feat[..., i].clamp(0, self.cfg.feat_vocab[i] - 1)) return self.drop(self.ln(x)) # --------------------------------------------------------------------------- # # Context encoder(可切換) # --------------------------------------------------------------------------- # class LocalAttention(nn.Module): """local-window self-attention + RoPE;|i-j| <= local_window 才可見。""" def __init__(self, cfg: TWLATRConfig): super().__init__() self.cfg = cfg self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model) self.out = nn.Linear(cfg.d_model, cfg.d_model) self.drop = nn.Dropout(cfg.dropout) def forward(self, x, attn_mask, rope): b, t, _ = x.shape h, d = self.cfg.n_heads, self.cfg.head_dim q, k, v = self.qkv(x).view(b, t, 3, h, d).permute(2, 0, 3, 1, 4).unbind(0) q, k = apply_rope(q, *rope), apply_rope(k, *rope) p = self.cfg.dropout if self.training else 0.0 y = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask, dropout_p=p) y = y.transpose(1, 2).reshape(b, t, self.cfg.d_model) return self.drop(self.out(y)) class ConvMixer(nn.Module): """dilated depthwise separable Conv1d:depthwise(k, dilation) → pointwise 擴張 → GELU → pointwise 還原。padding 位置先歸零,避免 pad 洩漏進感受野。""" def __init__(self, cfg: TWLATRConfig, dilation: int): super().__init__() d, k = cfg.d_model, cfg.conv_kernel pad = dilation * (k - 1) // 2 # 等長輸出 mid = d * cfg.conv_expansion self.dw = nn.Conv1d(d, d, k, padding=pad, dilation=dilation, groups=d) self.pw1 = nn.Conv1d(d, mid, 1) self.pw2 = nn.Conv1d(mid, d, 1) self.drop = nn.Dropout(cfg.dropout) def forward(self, x, pad_mask, rope=None): z = (x * pad_mask.unsqueeze(-1).to(x.dtype)).transpose(1, 2) z = self.pw2(F.gelu(self.pw1(self.dw(z)))) return self.drop(z.transpose(1, 2)) class EncoderBlock(nn.Module): """pre-LN:mixer(local attention 或 dilated conv)+ FFN。""" def __init__(self, cfg: TWLATRConfig, layer_idx: int): super().__init__() self.ln1 = nn.LayerNorm(cfg.d_model) if cfg.encoder == "local": self.mixer: nn.Module = LocalAttention(cfg) else: self.mixer = ConvMixer(cfg, cfg.dilation_at(layer_idx)) self.ln2 = nn.LayerNorm(cfg.d_model) self.ffn = nn.Sequential( nn.Linear(cfg.d_model, cfg.ffn_dim), nn.GELU(), nn.Dropout(cfg.dropout), nn.Linear(cfg.ffn_dim, cfg.d_model), nn.Dropout(cfg.dropout), ) def forward(self, x, mix_arg, rope=None): x = x + self.mixer(self.ln1(x), mix_arg, rope) x = x + self.ffn(self.ln2(x)) return x # --------------------------------------------------------------------------- # # TWLAT-R # --------------------------------------------------------------------------- # class TWLATR(nn.Module): def __init__(self, cfg: TWLATRConfig | None = None): super().__init__() self.cfg = cfg = cfg or TWLATRConfig() self.char_emb = CharEmbedding(cfg) self.text_feat = TextFeatures(cfg) self.layers = nn.ModuleList( EncoderBlock(cfg, i) for i in range(cfg.n_layers) ) self.enc_ln = nn.LayerNorm(cfg.d_model) # 候選側:mean-pool 後只有一個 LN + 一個線性投影,沒有任何 per-ID 參數 self.cand_ln = nn.LayerNorm(cfg.d_model) self.cand_proj = nn.Linear(cfg.d_model, cfg.d_model) # H-B 對照組:per-candidate learned lookup。未見過的候選只會取到 # 一列未訓練的隨機向量——這正是 V1 `cand_emb = nn.Embedding(4096, 192)` # 的行為,用來檢驗動態編碼是否真的帶來 held-out 優勢。 if cfg.candidate_encoder == "id": self.cand_id_emb = nn.Embedding(cfg.cand_id_vocab, cfg.d_model) score_in = 3 * cfg.d_model + cfg.cand_feat_dim self.score = nn.Sequential( nn.Linear(score_in, cfg.score_hidden), nn.GELU(), nn.Dropout(cfg.dropout), nn.Linear(cfg.score_hidden, 1), ) self.apply(self._init_weights) self._rope_cache: dict[Any, tuple[torch.Tensor, torch.Tensor]] = {} self._window_cache: dict[Any, torch.Tensor] = {} @staticmethod def _init_weights(m: nn.Module) -> None: if isinstance(m, nn.Linear): nn.init.normal_(m.weight, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) elif isinstance(m, nn.Embedding): nn.init.normal_(m.weight, std=0.02) elif isinstance(m, nn.Conv1d): nn.init.normal_(m.weight, std=0.02) if m.bias is not None: nn.init.zeros_(m.bias) # -- caches ------------------------------------------------------------ # def _rope(self, t: int, device, dtype): key = (t, str(device), dtype) if key not in self._rope_cache: self._rope_cache[key] = build_rope_cache( t, self.cfg.head_dim, self.cfg.rope_base, device, dtype ) return self._rope_cache[key] def _window_mask(self, t: int, device) -> torch.Tensor: key = (t, str(device)) if key not in self._window_cache: idx = torch.arange(t, device=device) self._window_cache[key] = (idx[:, None] - idx[None, :]).abs() <= self.cfg.local_window return self._window_cache[key] # -- 子步驟 ------------------------------------------------------------- # def encode_text(self, input_ids, feat, pad_mask) -> torch.Tensor: """[B,T] → [B,T,D]。""" t, device = input_ids.shape[1], input_ids.device x = self.text_feat(self.char_emb(input_ids), feat) if self.cfg.encoder == "local": ar = torch.arange(t, device=device) eye = ar[:, None] == ar[None, :] # 保留對角線,避免整列被遮而產生 NaN mask = ((self._window_mask(t, device) & pad_mask[:, None, :]) | eye).unsqueeze(1) rope = self._rope(t, device, x.dtype) for layer in self.layers: x = layer(x, mask, rope) else: for layer in self.layers: x = layer(x, pad_mask) return self.enc_ln(x) @staticmethod def span_pool(h: torch.Tensor, spans: torch.Tensor, valid: torch.Tensor) -> torch.Tensor: """對每個 proposal 的 [start,end) 做 mean-pool。h:[B,T,D] spans:[B,P,2] → [B,P,D]。""" t = h.shape[1] pos = torch.arange(t, device=h.device) start = spans[..., 0].clamp(0, t).unsqueeze(-1) end = spans[..., 1].clamp(0, t).unsqueeze(-1) w = ((pos >= start) & (pos < end) & valid).to(h.dtype) # [B,P,T] return torch.matmul(w, h) / w.sum(-1, keepdim=True).clamp(min=1.0) def encode_cands(self, cand_tok: torch.Tensor) -> torch.Tensor: """候選表面字串 → 向量。[B,P,C,L] → [B,P,C,D]。 走的是與文本同一份 char_emb,且只有 mean-pool + proj: 任何未見過的字串都能得到有限且有梯度的表徵(zero-shot 的前提)。 """ if self.cfg.candidate_encoder == "id": # 把字元序列雜湊成單一 id:同字串 → 同 id,不同字串 → 不同 id。 # 語意上等同 per-candidate 查表,且不需要重新產生資料。 mult = torch.tensor([1, 131, 131 ** 2, 131 ** 3, 131 ** 4, 131 ** 5], device=cand_tok.device, dtype=torch.long) mult = mult[: cand_tok.shape[-1]] cid = (cand_tok * mult).sum(-1) % self.cfg.cand_id_vocab return self.cand_id_emb(cid) valid = (cand_tok > 0).unsqueeze(-1) # id 0 = 右側 padding e = self.char_emb(cand_tok) e = e * valid.to(e.dtype) pooled = e.sum(-2) / valid.sum(-2).clamp(min=1).to(e.dtype) return self.cand_proj(self.cand_ln(pooled)) # -- forward ------------------------------------------------------------ # def forward(self, batch: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: """batch: input_ids [B,T] int64 原文字元 id feat [B,T,4] int64 script / 在 proposal span 內 / 保護段 / 詞界 pad_mask [B,T] bool (可省,缺省全 True) prop_spans [B,P,2] int64 proposal 的 [start,end) prop_mask [B,P] bool cand_tok [B,P,C,L] int64 候選表面字串(右側 0 padding) cand_mask [B,P,C] bool cand_feat [B,P,C,K] float 回傳 {"cand_logits": [B,P,C]},padding 候選為 -inf。 """ input_ids = batch["input_ids"] b, t = input_ids.shape device = input_ids.device pad_mask = batch.get("pad_mask") pad_mask = ( torch.ones(b, t, dtype=torch.bool, device=device) if pad_mask is None else pad_mask.bool() ) prop_mask = batch["prop_mask"].bool() cand_mask = batch["cand_mask"].bool() h_text = self.encode_text(input_ids, batch["feat"], pad_mask) h = self.span_pool( # [B,P,D] h_text, batch["prop_spans"], (pad_mask[:, None, :] & prop_mask[..., None]) ) e = self.encode_cands(batch["cand_tok"]) # [B,P,C,D] c = e.shape[2] h_exp = h.unsqueeze(2).expand(-1, -1, c, -1) z = torch.cat([h_exp, e, h_exp * e, batch["cand_feat"].to(e.dtype)], dim=-1) logits = self.score(z).squeeze(-1) # [B,P,C] return {"cand_logits": logits.masked_fill(~cand_mask, float("-inf"))} def compute_loss(self, outputs, batch): return compute_loss(outputs, batch) # --------------------------------------------------------------------------- # # Loss:L = L_cand + 0.3 · L_keep_bias # --------------------------------------------------------------------------- # def compute_loss( outputs: dict[str, torch.Tensor], batch: dict[str, torch.Tensor], weights: dict[str, float] | None = None, ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: """masked cross-entropy over candidates,外加 keep_bias 正則。 keep_bias:對 gold == 0(維持原樣)的 proposal,任何非 0 候選只要分數不比 候選 0 低 0.5 以上就受罰 —— 直接壓「改了不該改」(佔 V1 錯誤的 73%)。 """ w = {"cand": W_CAND, "keep_bias": W_KEEP_BIAS} if weights: w.update(weights) logits = outputs["cand_logits"] prop_mask = batch["prop_mask"].bool() cand_mask = batch["cand_mask"].bool() n_cands = logits.shape[-1] dtype = logits.dtype # padding 候選的 -inf 會污染算術:CE 用 finfo.min,hinge 用 0 並顯式遮罩 neg = torch.finfo(dtype).min safe = torch.where(cand_mask, logits, torch.full_like(logits, neg)) finite = torch.where(cand_mask, logits, torch.zeros_like(logits)) gold = batch["gold_cand"].clamp(0, n_cands - 1) prop_w = prop_mask.to(dtype) n_prop = prop_w.sum().clamp(min=1.0) # L_cand:每個 proposal 對候選集的 cross-entropy,padding proposal 不計 logp = torch.log_softmax(safe, dim=-1) nll = -logp.gather(-1, gold.unsqueeze(-1)).squeeze(-1) l_cand = (nll * prop_w).sum() / n_prop # L_keep_bias:gold 為「維持原樣」時,對每個非 0 候選各罰一次 hinge is_keep = (gold == KEEP_INDEX) & prop_mask s_keep = finite[..., KEEP_INDEX : KEEP_INDEX + 1] other = cand_mask & ( torch.arange(n_cands, device=logits.device)[None, None, :] != KEEP_INDEX ) hinge = F.relu(finite - s_keep + KEEP_MARGIN) * other.to(dtype) n_keep = is_keep.to(dtype).sum().clamp(min=1.0) l_keep = (hinge.sum(-1) * is_keep.to(dtype)).sum() / n_keep total = w["cand"] * l_cand + w["keep_bias"] * l_keep with torch.no_grad(): acc = ((safe.argmax(-1) == gold).to(dtype) * prop_w).sum() / n_prop return total, { "loss": total.detach(), "cand": l_cand.detach(), "keep_bias": l_keep.detach(), "acc": acc, "n_prop": n_prop.detach(), }