BorisTM/loss-guided-static-multi
08
1"""Exact batched matching primitives for a differentiable hard tokenizer."""2 3from __future__ import annotations4 5from dataclasses import dataclass6 7import torch8 9 10@dataclass(frozen=True)11class MatchingResult:12 """Partition statistics and deterministic MAP edges for a padded batch."""13 14 log_partition: torch.Tensor15 marginals: torch.Tensor16 map_edges: torch.Tensor17 18 19def _validate_edges(20 edge_scores: torch.Tensor, edge_mask: torch.Tensor21) -> tuple[torch.Tensor, torch.Tensor]:22 if edge_scores.ndim != 2 or edge_mask.ndim != 2:23 raise ValueError("edge_scores and edge_mask must both have shape [batch, edges]")24 if edge_scores.shape != edge_mask.shape:25 raise ValueError("edge_scores and edge_mask must have identical shapes")26 if not edge_scores.is_floating_point():27 raise TypeError("edge_scores must be floating point")28 if edge_mask.dtype is not torch.bool:29 raise TypeError("edge_mask must be boolean")30 if edge_scores.device != edge_mask.device:31 raise ValueError("edge_scores and edge_mask must use the same device")32 if bool((~torch.isfinite(edge_scores) & edge_mask).any()):33 raise ValueError("valid edge scores must be finite")34 scores = torch.where(edge_mask, edge_scores.float(), torch.zeros_like(edge_scores.float()))35 return scores, edge_mask36 37 38def _prefix_log_partitions(39 scores: torch.Tensor, mask: torch.Tensor40) -> list[torch.Tensor]:41 batch = scores.shape[0]42 zero = torch.zeros(batch, dtype=torch.float32, device=scores.device)43 prefix = [zero, zero]44 for edge in range(scores.shape[1]):45 separate = prefix[-1]46 merged = prefix[-2] + scores[:, edge]47 prefix.append(48 torch.where(mask[:, edge], torch.logaddexp(separate, merged), separate)49 )50 return prefix51 52 53def _suffix_log_partitions(54 scores: torch.Tensor, mask: torch.Tensor55) -> list[torch.Tensor]:56 batch, edges = scores.shape57 zero = torch.zeros(batch, dtype=torch.float32, device=scores.device)58 suffix = [zero for _ in range(edges + 2)]59 for edge in range(edges - 1, -1, -1):60 separate = suffix[edge + 1]61 merged = scores[:, edge] + suffix[edge + 2]62 suffix[edge] = torch.where(63 mask[:, edge], torch.logaddexp(separate, merged), separate64 )65 return suffix66 67 68def _map_matching(scores: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:69 batch, edges = scores.shape70 if edges == 0:71 return torch.zeros_like(mask)72 73 zero = torch.zeros(batch, dtype=torch.float32, device=scores.device)74 best = [zero, zero]75 take_by_edge = torch.zeros_like(mask)76 for edge in range(edges):77 separate = best[-1]78 merged = best[-2] + scores[:, edge]79 take = mask[:, edge] & (merged > separate)80 best.append(torch.where(take, merged, separate))81 take_by_edge[:, edge] = take82 83 selected = torch.zeros_like(mask)84 vertices = torch.full(85 (batch,), edges + 1, dtype=torch.long, device=scores.device86 )87 rows = torch.arange(batch, device=scores.device)88 for _ in range(edges + 1):89 active = vertices >= 290 edge = (vertices - 2).clamp(min=0, max=edges - 1)91 take = active & take_by_edge[rows, edge]92 selected[rows, edge] |= take93 vertices = vertices - torch.where(take, 2, 1) * active.to(torch.long)94 return selected95 96 97def batched_matching(98 edge_scores: torch.Tensor,99 edge_mask: torch.Tensor,100) -> MatchingResult:101 """Solve independent monomer-dimer CRFs over a padded sentence batch.102 103 Invalid edges act as fixed token boundaries. Every probabilistic recurrence104 runs in float32 even when the caller is inside BF16 autocast.105 """106 107 scores, mask = _validate_edges(edge_scores, edge_mask)108 prefix = _prefix_log_partitions(scores, mask)109 log_partition = prefix[-1]110 111 if scores.shape[1] == 0:112 marginals = torch.empty_like(scores)113 else:114 suffix = _suffix_log_partitions(scores, mask)115 log_marginals = torch.stack(116 [117 prefix[edge]118 + scores[:, edge]119 + suffix[edge + 2]120 - log_partition121 for edge in range(scores.shape[1])122 ],123 dim=1,124 )125 marginals = torch.where(mask, torch.exp(log_marginals), torch.zeros_like(scores))126 127 map_edges = _map_matching(scores, mask)128 if not bool(torch.isfinite(log_partition).all()):129 raise FloatingPointError("nonfinite matching log partition")130 if not bool(torch.isfinite(marginals).all()):131 raise FloatingPointError("nonfinite matching marginals")132 return MatchingResult(log_partition, marginals, map_edges)133 134 135def structured_straight_through(136 marginals: torch.Tensor,137 map_edges: torch.Tensor,138 *,139 dtype: torch.dtype,140) -> torch.Tensor:141 """Return hard MAP values whose gradient follows exact edge marginals."""142 143 if marginals.shape != map_edges.shape:144 raise ValueError("marginals and map_edges must have identical shapes")145 if map_edges.dtype is not torch.bool:146 raise TypeError("map_edges must be boolean")147 if not dtype.is_floating_point:148 raise TypeError("straight-through dtype must be floating point")149 soft = marginals.to(dtype=dtype)150 hard = map_edges.to(dtype=dtype)151 return soft + (hard - soft).detach()152 