"""TWLAT V3:lattice 上的雙 pass cloze 模型。 任務:對 lattice 的每條邊,從 confusion group 的正規候選集中預測 「臺灣書寫者在這個語境會寫哪個形式」。 clean pass(汙染文本原樣)──→ h_clean ─┐ 表面形式證據 masked pass(站點收合為 MASK)→ h_m ───┤ 無洩漏語境證據 候選(共用 char_emb 動態編碼)→ e_c ───┼→ score MLP → [B,S,C] 字典特徵(64 維 lattice 特徵)─────────┘ 與 V2 的差異: 1. 預訓練時 observed 相依特徵歸零(collate 控制),模型無法走 「相信表面」捷徑;finetune 才學習把表面 prior 併進來。 2. doc 向量在中層注入:領域相依詞(程序/數據/介面)需要全文域推斷。 3. MLM 輔助頭(tied embedding)維持表徵品質。 參數量(d256 / 8 層 / 2 attn)≈ 8.7M,遠低於 16M 上限(D-04)。 """ from __future__ import annotations import dataclasses from dataclasses import dataclass from typing import Any import torch import torch.nn as nn import torch.nn.functional as F from twlat.model_r import (CharEmbedding, ConvMixer, LocalAttention, TextFeatures, build_rope_cache) @dataclass class TWLATV3Config: d_model: int = 256 n_heads: int = 4 ffn_dim: int = 1024 dropout: float = 0.1 n_layers: int = 8 attn_layers: tuple[int, ...] = (3, 7) # 這些層用 local attention,其餘 TCN local_window: int = 64 conv_kernel: int = 5 conv_dilations: tuple[int, ...] = (1, 2, 4, 8, 16, 32, 64, 128) conv_expansion: int = 2 doc_layer: int = 4 # 此層之前注入 doc 向量 han_vocab: int = 4096 # 0=PAD 1=UNK 2=MASK n_hash: int = 2 hash_buckets: int = 2048 feat_dim: int = 4 # script(含 MASK=4)/span 內/保護段/詞界 feat_vocab: tuple[int, ...] = (16, 4, 4, 4) s_max: int = 128 c_max: int = 8 cand_len: int = 8 cand_feat_dim: int = 64 score_hidden: int = 512 seq_len: int = 512 rope_base: float = 10000.0 # loss mlm_weight: float = 0.1 nomask_weight: float = 0.1 # 不可遮罩站點的 loss 權重 keep_margin: float = 0.5 # finetune 階段的 keep hinge def __post_init__(self): assert self.d_model % self.n_heads == 0 and self.d_model % 2 == 0 assert self.conv_kernel % 2 == 1 @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)] class V3Block(nn.Module): """pre-LN block;mixer 依層選 local attention 或 dilated conv。""" def __init__(self, cfg: TWLATV3Config, layer_idx: int): super().__init__() self.is_attn = layer_idx in cfg.attn_layers self.ln1 = nn.LayerNorm(cfg.d_model) if self.is_attn: 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, attn_mask, pad_mask, rope): if self.is_attn: x = x + self.mixer(self.ln1(x), attn_mask, rope) else: x = x + self.mixer(self.ln1(x), pad_mask) return x + self.ffn(self.ln2(x)) class TWLATV3(nn.Module): def __init__(self, cfg: TWLATV3Config | None = None): super().__init__() self.cfg = cfg = cfg or TWLATV3Config() self.char_emb = CharEmbedding(cfg) self.text_feat = TextFeatures(cfg) self.layers = nn.ModuleList(V3Block(cfg, i) for i in range(cfg.n_layers)) self.enc_ln = nn.LayerNorm(cfg.d_model) self.doc_mlp = nn.Sequential( nn.Linear(cfg.d_model, cfg.d_model), nn.GELU(), nn.Linear(cfg.d_model, cfg.d_model)) self.cand_ln = nn.LayerNorm(cfg.d_model) self.cand_proj = nn.Linear(cfg.d_model, cfg.d_model) score_in = 5 * 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] = {} self._window_cache: dict[Any, torch.Tensor] = {} @staticmethod def _init_weights(m): if isinstance(m, (nn.Linear, nn.Conv1d)): 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) def _rope(self, t, 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 _win_mask(self, t, device): 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(self, ids, feat, pad_mask) -> torch.Tensor: """[B,T] → [B,T,D],中層注入 doc mean-pool 向量(域推斷通道)。""" t, device = ids.shape[1], ids.device x = self.text_feat(self.char_emb(ids), feat) ar = torch.arange(t, device=device) eye = ar[:, None] == ar[None, :] attn_mask = ((self._win_mask(t, device) & pad_mask[:, None, :]) | eye ).unsqueeze(1) rope = self._rope(t, device, x.dtype) pw = pad_mask.unsqueeze(-1).to(x.dtype) for i, layer in enumerate(self.layers): if i == self.cfg.doc_layer: doc = (x * pw).sum(1) / pw.sum(1).clamp(min=1.0) x = x + self.doc_mlp(doc).unsqueeze(1) x = layer(x, attn_mask, pad_mask, rope) return self.enc_ln(x) @staticmethod def span_pool(h, spans, valid): 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) return torch.matmul(w, h) / w.sum(-1, keepdim=True).clamp(min=1.0) def encode_cands(self, cand_tok) -> torch.Tensor: """[B,S,C,L] → [B,S,C,D];共用 char_emb,零 per-ID 參數(熱更新前提)。""" valid = (cand_tok > 0).unsqueeze(-1) e = self.char_emb(cand_tok) * valid.to(self.char_emb.han.weight.dtype) pooled = e.sum(-2) / valid.sum(-2).clamp(min=1).to(e.dtype) return self.cand_proj(self.cand_ln(pooled)) def mlm_logits(self, h) -> torch.Tensor: """tied 到 han embedding(只覆蓋常用字表)。""" return h @ self.char_emb.han.weight.t() def forward(self, batch: dict[str, torch.Tensor]) -> dict[str, torch.Tensor]: """batch 欄位: ids/feat/pad clean 序列 [B,T]… mids/mfeat/mpad masked 序列 [B,Tm]… c_span/m_span [B,S,2] 兩序列座標 site_mask [B,S] cand_tok [B,S,C,L] cand_mask/cand_kill [B,S,C] cand_feat [B,S,C,K] (訓練另有 gold/site_w/mlm_pos/mlm_gold) """ site_mask = batch["site_mask"].bool() cand_ok = batch["cand_mask"].bool() & ~batch["cand_kill"].bool() h_c = self.encode(batch["ids"], batch["feat"], batch["pad"].bool()) h_m = self.encode(batch["mids"], batch["mfeat"], batch["mpad"].bool()) vc = batch["pad"].bool()[:, None, :] & site_mask[..., None] vm = batch["mpad"].bool()[:, None, :] & site_mask[..., None] hc = self.span_pool(h_c, batch["c_span"], vc) # [B,S,D] hm = self.span_pool(h_m, batch["m_span"], vm) e = self.encode_cands(batch["cand_tok"]) # [B,S,C,D] c = e.shape[2] hce = hc.unsqueeze(2).expand(-1, -1, c, -1) hme = hm.unsqueeze(2).expand(-1, -1, c, -1) z = torch.cat([hme, hce, e, hme * e, hce * e, batch["cand_feat"].to(e.dtype)], dim=-1) logits = self.score(z).squeeze(-1) return {"cand_logits": logits.masked_fill(~cand_ok, float("-inf")), "h_m": h_m} def compute_loss_v3(model: TWLATV3, out, batch, phase: str = "pretrain"): cfg = model.cfg logits = out["cand_logits"] dtype = logits.dtype gold = batch["gold"].clamp(min=0) # gold 候選被硬過濾砍掉的位點(例外詞/positional 與真實用法衝突): # 模型無從答對,排除於 loss——這是硬過濾的固有代價,由 gold_killed 計數監控 selectable = batch["cand_mask"].bool() & ~batch["cand_kill"].bool() gold_ok = selectable.gather(-1, gold.unsqueeze(-1)).squeeze(-1) site_mask = batch["site_mask"].bool() & (batch["gold"] >= 0) & gold_ok w = batch["site_w"].to(dtype) * site_mask.to(dtype) n = w.sum().clamp(min=1.0) neg = torch.finfo(dtype).min safe = torch.where(torch.isinf(logits), torch.full_like(logits, neg), logits) logp = torch.log_softmax(safe, -1) nll = -logp.gather(-1, gold.unsqueeze(-1)).squeeze(-1) l_cloze = (nll * w).sum() / n total = l_cloze parts = {"cloze": l_cloze.detach()} if "mlm_pos" in batch and batch["mlm_pos"].any(): ml = model.mlm_logits(out["h_m"]) pos = batch["mlm_pos"].bool() l_mlm = F.cross_entropy(ml[pos], batch["mlm_gold"][pos].clamp(min=0)) total = total + cfg.mlm_weight * l_mlm parts["mlm"] = l_mlm.detach() if phase == "finetune": # keep hinge:gold==observed 時,其他候選高過 s_obs−margin 即受罰 obs = batch["obs"].long().clamp(min=0) is_keep = site_mask & (batch["gold"] == batch["obs"]) finite = torch.where(torch.isinf(logits), torch.zeros_like(logits), logits) s_obs = finite.gather(-1, obs.unsqueeze(-1)) others = batch["cand_mask"].bool() & ~batch["cand_kill"].bool() & \ (torch.arange(logits.shape[-1], device=logits.device)[None, None, :] != obs.unsqueeze(-1)) hinge = F.relu(finite - s_obs + cfg.keep_margin) * others.to(dtype) nk = is_keep.to(dtype).sum().clamp(min=1.0) l_keep = (hinge.sum(-1) * is_keep.to(dtype)).sum() / nk total = total + 0.15 * l_keep parts["keep"] = l_keep.detach() with torch.no_grad(): pred = safe.argmax(-1) ok = (pred == gold) & site_mask keepm = site_mask & (batch["gold"] == batch["obs"]) chgm = site_mask & (batch["gold"] != batch["obs"]) parts.update( acc=ok.float().sum() / site_mask.float().sum().clamp(min=1.0), keep_acc=(ok & keepm).float().sum() / keepm.float().sum().clamp(min=1.0), chg_acc=(ok & chgm).float().sum() / chgm.float().sum().clamp(min=1.0), n_sites=site_mask.float().sum(), gold_killed=(batch["site_mask"].bool() & (batch["gold"] >= 0) & ~gold_ok).float().sum()) parts["loss"] = total.detach() return total, parts def make_config(**over) -> TWLATV3Config: fields = {f.name for f in dataclasses.fields(TWLATV3Config)} return TWLATV3Config(**{k: (tuple(v) if isinstance(v, list) else v) for k, v in over.items() if k in fields})