CoolFace
Modelpublic

BorisTM/loss-guided-static-multi

sourceHugging Faceapache-2.0updated 23h agoView on Hugging Face
0likes8downloads
language_conditioned.py166 linesDownload Raw Back to root
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