CoolFace
Modelpublic

BorisTM/loss-guided-static-multi

sourceHugging Faceapache-2.0updated 23h agoView on Hugging Face
0likes8downloads
structured_tokenizer.py152 linesDownload Raw Back to root
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