loss-guided-static-multi / structured_tokenizer.py
BorisTM's picture
Add files using upload-large-folder tool
309d3a9 verified
Raw History Blame Contribute Delete
5.38 kB
"""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()