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: 5,376 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 | """Exact batched matching primitives for a differentiable hard tokenizer."""
from __future__ import annotations
from dataclasses import dataclass
import torch
@dataclass(frozen=True)
class MatchingResult:
"""Partition statistics and deterministic MAP edges for a padded batch."""
log_partition: torch.Tensor
marginals: torch.Tensor
map_edges: torch.Tensor
def _validate_edges(
edge_scores: torch.Tensor, edge_mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
if edge_scores.ndim != 2 or edge_mask.ndim != 2:
raise ValueError("edge_scores and edge_mask must both have shape [batch, edges]")
if edge_scores.shape != edge_mask.shape:
raise ValueError("edge_scores and edge_mask must have identical shapes")
if not edge_scores.is_floating_point():
raise TypeError("edge_scores must be floating point")
if edge_mask.dtype is not torch.bool:
raise TypeError("edge_mask must be boolean")
if edge_scores.device != edge_mask.device:
raise ValueError("edge_scores and edge_mask must use the same device")
if bool((~torch.isfinite(edge_scores) & edge_mask).any()):
raise ValueError("valid edge scores must be finite")
scores = torch.where(edge_mask, edge_scores.float(), torch.zeros_like(edge_scores.float()))
return scores, edge_mask
def _prefix_log_partitions(
scores: torch.Tensor, mask: torch.Tensor
) -> list[torch.Tensor]:
batch = scores.shape[0]
zero = torch.zeros(batch, dtype=torch.float32, device=scores.device)
prefix = [zero, zero]
for edge in range(scores.shape[1]):
separate = prefix[-1]
merged = prefix[-2] + scores[:, edge]
prefix.append(
torch.where(mask[:, edge], torch.logaddexp(separate, merged), separate)
)
return prefix
def _suffix_log_partitions(
scores: torch.Tensor, mask: torch.Tensor
) -> list[torch.Tensor]:
batch, edges = scores.shape
zero = torch.zeros(batch, dtype=torch.float32, device=scores.device)
suffix = [zero for _ in range(edges + 2)]
for edge in range(edges - 1, -1, -1):
separate = suffix[edge + 1]
merged = scores[:, edge] + suffix[edge + 2]
suffix[edge] = torch.where(
mask[:, edge], torch.logaddexp(separate, merged), separate
)
return suffix
def _map_matching(scores: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
batch, edges = scores.shape
if edges == 0:
return torch.zeros_like(mask)
zero = torch.zeros(batch, dtype=torch.float32, device=scores.device)
best = [zero, zero]
take_by_edge = torch.zeros_like(mask)
for edge in range(edges):
separate = best[-1]
merged = best[-2] + scores[:, edge]
take = mask[:, edge] & (merged > separate)
best.append(torch.where(take, merged, separate))
take_by_edge[:, edge] = take
selected = torch.zeros_like(mask)
vertices = torch.full(
(batch,), edges + 1, dtype=torch.long, device=scores.device
)
rows = torch.arange(batch, device=scores.device)
for _ in range(edges + 1):
active = vertices >= 2
edge = (vertices - 2).clamp(min=0, max=edges - 1)
take = active & take_by_edge[rows, edge]
selected[rows, edge] |= take
vertices = vertices - torch.where(take, 2, 1) * active.to(torch.long)
return selected
def batched_matching(
edge_scores: torch.Tensor,
edge_mask: torch.Tensor,
) -> MatchingResult:
"""Solve independent monomer-dimer CRFs over a padded sentence batch.
Invalid edges act as fixed token boundaries. Every probabilistic recurrence
runs in float32 even when the caller is inside BF16 autocast.
"""
scores, mask = _validate_edges(edge_scores, edge_mask)
prefix = _prefix_log_partitions(scores, mask)
log_partition = prefix[-1]
if scores.shape[1] == 0:
marginals = torch.empty_like(scores)
else:
suffix = _suffix_log_partitions(scores, mask)
log_marginals = torch.stack(
[
prefix[edge]
+ scores[:, edge]
+ suffix[edge + 2]
- log_partition
for edge in range(scores.shape[1])
],
dim=1,
)
marginals = torch.where(mask, torch.exp(log_marginals), torch.zeros_like(scores))
map_edges = _map_matching(scores, mask)
if not bool(torch.isfinite(log_partition).all()):
raise FloatingPointError("nonfinite matching log partition")
if not bool(torch.isfinite(marginals).all()):
raise FloatingPointError("nonfinite matching marginals")
return MatchingResult(log_partition, marginals, map_edges)
def structured_straight_through(
marginals: torch.Tensor,
map_edges: torch.Tensor,
*,
dtype: torch.dtype,
) -> torch.Tensor:
"""Return hard MAP values whose gradient follows exact edge marginals."""
if marginals.shape != map_edges.shape:
raise ValueError("marginals and map_edges must have identical shapes")
if map_edges.dtype is not torch.bool:
raise TypeError("map_edges must be boolean")
if not dtype.is_floating_point:
raise TypeError("straight-through dtype must be floating point")
soft = marginals.to(dtype=dtype)
hard = map_edges.to(dtype=dtype)
return soft + (hard - soft).detach()
|