BorisTM/loss-guided-static-multi
08
1"""Language-conditioned static embedding.2 3The language is carried as a marker token prepended to each text, so the two4sides of a parallel pair can carry different language identities without5changing the trainer's two-column dataset schema.6"""7 8from __future__ import annotations9 10import torch11from torch import nn12 13MODES = (14 "none", "marker", "centroid", "diag", "lowrank", "bilinear", "senses",15 "hard", "gate", "gate_senses", "split", "ngram", "ngram3", "wordpool",16 "ngram_gate", "fuse", "vocabmoe", "collapse", "collapse_shared",17 "collapse_dynamic", "collapse_reclaimed", "collapse_online",18 "collapse_online_compiled", "collapse_recursive", "idfpool",19)20 21# Language-specific pooling changes the relative numerator composition. One22# shared vocabulary row cannot represent a different weight in every language;23# the positive weighted-mean denominator itself cancels under cosine scoring.24GATING_MODES = ("gate", "gate_senses")25 26 27class LanguageConditioner(nn.Module):28 """The conditioning head. Owns every parameter the base table does not."""29 30 def __init__(31 self,32 *,33 mode: str,34 dim: int,35 vocab_size: int,36 n_languages: int,37 code_dim: int = 16,38 n_senses: int = 64,39 rank: int = 8,40 ) -> None:41 super().__init__()42 if mode not in MODES:43 raise ValueError(f"unknown mode {mode!r}, expected one of {MODES}")44 self.mode = mode45 self.dim = dim46 self.n_languages = n_languages47 self.code_dim = code_dim48 self.n_senses = n_senses49 self.rank = rank50 51 if mode == "centroid":52 self.centroid = nn.Parameter(torch.zeros(n_languages, dim))53 elif mode == "diag":54 self.scale = nn.Parameter(torch.ones(n_languages, dim))55 elif mode == "lowrank":56 self.left = nn.Parameter(torch.randn(n_languages, dim, rank) * 0.02)57 self.right = nn.Parameter(torch.zeros(n_languages, rank, dim))58 elif mode == "bilinear":59 self.code = nn.EmbeddingBag(vocab_size, code_dim, mode="mean")60 nn.init.normal_(self.code.weight, std=0.02)61 self.readout = nn.Parameter(torch.zeros(n_languages, code_dim, dim))62 if mode in ("gate", "gate_senses"):63 self.gate_code = nn.Embedding(vocab_size, code_dim)64 nn.init.normal_(self.gate_code.weight, std=0.02)65 self.gate_language = nn.Parameter(torch.zeros(n_languages, code_dim))66 if mode in ("senses", "hard", "gate_senses"):67 self.code = nn.Embedding(vocab_size, code_dim)68 nn.init.normal_(self.code.weight, std=0.02)69 self.language = nn.Parameter(torch.ones(n_languages, code_dim))70 self.probe = nn.Parameter(torch.randn(n_senses, code_dim) * 0.02)71 self.senses = nn.Parameter(torch.zeros(n_senses, dim))72 73 def extra_repr(self) -> str:74 return (f"mode={self.mode}, languages={self.n_languages}, "75 f"code_dim={self.code_dim}, senses={self.n_senses}, rank={self.rank}")76 77 def token_gate(self, token_ids: torch.Tensor, lang_per_token: torch.Tensor) -> torch.Tensor:78 """Multiplicative weight per token, centred on 1 at initialisation."""79 score = (self.gate_code(token_ids) * self.gate_language[lang_per_token]).sum(-1)80 return 2.0 * torch.sigmoid(score)81 82 def sense_weights(self, token_ids: torch.Tensor, lang_per_token: torch.Tensor) -> torch.Tensor:83 """Alpha over the sense dictionary, per token, given its language."""84 a = self.code(token_ids)85 v = self.language[lang_per_token]86 logits = (a * v) @ self.probe.T87 alpha = torch.softmax(logits, dim=-1)88 if self.mode == "hard":89 index = alpha.argmax(dim=-1, keepdim=True)90 onehot = torch.zeros_like(alpha).scatter_(-1, index, 1.0)91 alpha = onehot + alpha - alpha.detach()92 return alpha93 94 def forward(95 self,96 pooled: torch.Tensor,97 language: torch.Tensor,98 token_ids: torch.Tensor,99 segment: torch.Tensor,100 lengths: torch.Tensor,101 ) -> torch.Tensor:102 mode = self.mode103 if mode in ("none", "marker"):104 return pooled105 if mode == "centroid":106 return pooled - self.centroid[language]107 if mode == "diag":108 return pooled * self.scale[language]109 if mode == "lowrank":110 left = self.left[language]111 right = self.right[language]112 latent = torch.einsum("bd,bdr->br", pooled, left)113 return pooled + torch.einsum("br,brd->bd", latent, right)114 if mode == "bilinear":115 offsets = torch.cat([116 torch.zeros(1, dtype=torch.long, device=token_ids.device),117 lengths.cumsum(0)[:-1],118 ])119 mean_code = self.code(token_ids, offsets)120 return pooled + torch.einsum("bk,bkd->bd", mean_code, self.readout[language])121 if mode == "gate":122 return pooled123 124 alpha = self.sense_weights(token_ids, language[segment])125 summed = torch.zeros(pooled.shape[0], self.n_senses,126 dtype=alpha.dtype, device=alpha.device)127 summed.index_add_(0, segment, alpha)128 mean_alpha = summed / lengths.clamp(min=1).unsqueeze(1).to(summed.dtype)129 return pooled + mean_alpha @ self.senses130 131 132def split_markers(133 input_ids: torch.Tensor,134 offsets: torch.Tensor,135 marker_lookup: torch.Tensor,136 keep_marker: bool,137) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:138 """Peel the leading language marker off every sequence.139 140 Returns content ids, sentence indices, content lengths, per-sentence141 languages, and offsets for the content-only token stream.142 """143 total = input_ids.numel()144 batch = offsets.numel()145 marker_ids = input_ids[offsets]146 language = marker_lookup[marker_ids]147 148 ends = torch.cat([offsets[1:], torch.tensor([total], device=offsets.device)])149 full_lengths = ends - offsets150 if keep_marker:151 segment = torch.repeat_interleave(152 torch.arange(batch, device=offsets.device), full_lengths)153 return input_ids, segment, full_lengths, language, offsets154 155 keep = torch.ones(total, dtype=torch.bool, device=input_ids.device)156 keep[offsets] = False157 content = input_ids[keep]158 lengths = (full_lengths - 1).clamp(min=0)159 segment = torch.repeat_interleave(160 torch.arange(batch, device=offsets.device), lengths)161 new_offsets = torch.cat([162 torch.zeros(1, dtype=torch.long, device=offsets.device),163 lengths.cumsum(0)[:-1],164 ])165 return content, segment, lengths, language, new_offsets166 