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
File size: 6,761 Bytes
309d3a9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 | """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
|