Instructions to use BorisTM/loss-guided-static-multi with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use BorisTM/loss-guided-static-multi with sentence-transformers:
from sentence_transformers import SentenceTransformer model = SentenceTransformer("BorisTM/loss-guided-static-multi") sentences = [ "The weather is lovely today.", "It's so sunny outside!", "He drove to the stadium." ] embeddings = model.encode(sentences) similarities = model.similarity(embeddings, embeddings) print(similarities.shape) # [3, 3] - Notebooks
- Google Colab
- Kaggle
Download language_conditioned.py from BorisTM/loss-guided-static-multi: direct link, hf CLI and curl.
- Browser
- Download file 6.76 kB
-
https://huggingface.co/BorisTM/loss-guided-static-multi/resolve/main/language_conditioned.py
- Command line
-
hf download hf://BorisTM/loss-guided-static-multi/language_conditioned.py
-
curl -L -o language_conditioned.py https://huggingface.co/BorisTM/loss-guided-static-multi/resolve/main/language_conditioned.py
6.76 kB
| """Language-conditioned static embedding. | |
| The language is carried as a marker token prepended to each text, so the two | |
| sides of a parallel pair can carry different language identities without | |
| changing the trainer's two-column dataset schema. | |
| """ | |
| from __future__ import annotations | |
| import torch | |
| from torch import nn | |
| MODES = ( | |
| "none", "marker", "centroid", "diag", "lowrank", "bilinear", "senses", | |
| "hard", "gate", "gate_senses", "split", "ngram", "ngram3", "wordpool", | |
| "ngram_gate", "fuse", "vocabmoe", "collapse", "collapse_shared", | |
| "collapse_dynamic", "collapse_reclaimed", "collapse_online", | |
| "collapse_online_compiled", "collapse_recursive", "idfpool", | |
| ) | |
| # Language-specific pooling changes the relative numerator composition. One | |
| # shared vocabulary row cannot represent a different weight in every language; | |
| # the positive weighted-mean denominator itself cancels under cosine scoring. | |
| GATING_MODES = ("gate", "gate_senses") | |
| class LanguageConditioner(nn.Module): | |
| """The conditioning head. Owns every parameter the base table does not.""" | |
| def __init__( | |
| self, | |
| *, | |
| mode: str, | |
| dim: int, | |
| vocab_size: int, | |
| n_languages: int, | |
| code_dim: int = 16, | |
| n_senses: int = 64, | |
| rank: int = 8, | |
| ) -> None: | |
| super().__init__() | |
| if mode not in MODES: | |
| raise ValueError(f"unknown mode {mode!r}, expected one of {MODES}") | |
| self.mode = mode | |
| self.dim = dim | |
| self.n_languages = n_languages | |
| self.code_dim = code_dim | |
| self.n_senses = n_senses | |
| self.rank = rank | |
| if mode == "centroid": | |
| self.centroid = nn.Parameter(torch.zeros(n_languages, dim)) | |
| elif mode == "diag": | |
| self.scale = nn.Parameter(torch.ones(n_languages, dim)) | |
| elif mode == "lowrank": | |
| self.left = nn.Parameter(torch.randn(n_languages, dim, rank) * 0.02) | |
| self.right = nn.Parameter(torch.zeros(n_languages, rank, dim)) | |
| elif mode == "bilinear": | |
| self.code = nn.EmbeddingBag(vocab_size, code_dim, mode="mean") | |
| nn.init.normal_(self.code.weight, std=0.02) | |
| self.readout = nn.Parameter(torch.zeros(n_languages, code_dim, dim)) | |
| if mode in ("gate", "gate_senses"): | |
| self.gate_code = nn.Embedding(vocab_size, code_dim) | |
| nn.init.normal_(self.gate_code.weight, std=0.02) | |
| self.gate_language = nn.Parameter(torch.zeros(n_languages, code_dim)) | |
| if mode in ("senses", "hard", "gate_senses"): | |
| self.code = nn.Embedding(vocab_size, code_dim) | |
| nn.init.normal_(self.code.weight, std=0.02) | |
| self.language = nn.Parameter(torch.ones(n_languages, code_dim)) | |
| self.probe = nn.Parameter(torch.randn(n_senses, code_dim) * 0.02) | |
| self.senses = nn.Parameter(torch.zeros(n_senses, dim)) | |
| def extra_repr(self) -> str: | |
| return (f"mode={self.mode}, languages={self.n_languages}, " | |
| f"code_dim={self.code_dim}, senses={self.n_senses}, rank={self.rank}") | |
| def token_gate(self, token_ids: torch.Tensor, lang_per_token: torch.Tensor) -> torch.Tensor: | |
| """Multiplicative weight per token, centred on 1 at initialisation.""" | |
| score = (self.gate_code(token_ids) * self.gate_language[lang_per_token]).sum(-1) | |
| return 2.0 * torch.sigmoid(score) | |
| def sense_weights(self, token_ids: torch.Tensor, lang_per_token: torch.Tensor) -> torch.Tensor: | |
| """Alpha over the sense dictionary, per token, given its language.""" | |
| a = self.code(token_ids) | |
| v = self.language[lang_per_token] | |
| logits = (a * v) @ self.probe.T | |
| alpha = torch.softmax(logits, dim=-1) | |
| if self.mode == "hard": | |
| index = alpha.argmax(dim=-1, keepdim=True) | |
| onehot = torch.zeros_like(alpha).scatter_(-1, index, 1.0) | |
| alpha = onehot + alpha - alpha.detach() | |
| return alpha | |
| def forward( | |
| self, | |
| pooled: torch.Tensor, | |
| language: torch.Tensor, | |
| token_ids: torch.Tensor, | |
| segment: torch.Tensor, | |
| lengths: torch.Tensor, | |
| ) -> torch.Tensor: | |
| mode = self.mode | |
| if mode in ("none", "marker"): | |
| return pooled | |
| if mode == "centroid": | |
| return pooled - self.centroid[language] | |
| if mode == "diag": | |
| return pooled * self.scale[language] | |
| if mode == "lowrank": | |
| left = self.left[language] | |
| right = self.right[language] | |
| latent = torch.einsum("bd,bdr->br", pooled, left) | |
| return pooled + torch.einsum("br,brd->bd", latent, right) | |
| if mode == "bilinear": | |
| offsets = torch.cat([ | |
| torch.zeros(1, dtype=torch.long, device=token_ids.device), | |
| lengths.cumsum(0)[:-1], | |
| ]) | |
| mean_code = self.code(token_ids, offsets) | |
| return pooled + torch.einsum("bk,bkd->bd", mean_code, self.readout[language]) | |
| if mode == "gate": | |
| return pooled | |
| alpha = self.sense_weights(token_ids, language[segment]) | |
| summed = torch.zeros(pooled.shape[0], self.n_senses, | |
| dtype=alpha.dtype, device=alpha.device) | |
| summed.index_add_(0, segment, alpha) | |
| mean_alpha = summed / lengths.clamp(min=1).unsqueeze(1).to(summed.dtype) | |
| return pooled + mean_alpha @ self.senses | |
| def split_markers( | |
| input_ids: torch.Tensor, | |
| offsets: torch.Tensor, | |
| marker_lookup: torch.Tensor, | |
| keep_marker: bool, | |
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: | |
| """Peel the leading language marker off every sequence. | |
| Returns content ids, sentence indices, content lengths, per-sentence | |
| languages, and offsets for the content-only token stream. | |
| """ | |
| total = input_ids.numel() | |
| batch = offsets.numel() | |
| marker_ids = input_ids[offsets] | |
| language = marker_lookup[marker_ids] | |
| ends = torch.cat([offsets[1:], torch.tensor([total], device=offsets.device)]) | |
| full_lengths = ends - offsets | |
| if keep_marker: | |
| segment = torch.repeat_interleave( | |
| torch.arange(batch, device=offsets.device), full_lengths) | |
| return input_ids, segment, full_lengths, language, offsets | |
| keep = torch.ones(total, dtype=torch.bool, device=input_ids.device) | |
| keep[offsets] = False | |
| content = input_ids[keep] | |
| lengths = (full_lengths - 1).clamp(min=0) | |
| segment = torch.repeat_interleave( | |
| torch.arange(batch, device=offsets.device), lengths) | |
| new_offsets = torch.cat([ | |
| torch.zeros(1, dtype=torch.long, device=offsets.device), | |
| lengths.cumsum(0)[:-1], | |
| ]) | |
| return content, segment, lengths, language, new_offsets | |