decosa-cell-pointer-modernbert-base / modeling_cellpointer.py
sipratt's picture
decosa-cell-pointer-modernbert-base: weights, card, eval summary (Apache-2.0)
2769c0b verified
Raw History Blame Contribute Delete
2.86 kB
"""The cell-pointer: an encoder over [tagged sentences] [SEP] [tables written cell by cell], and small heads.
pointer per (tagged number, cell) pair: is this the cell the sentence says the number reports? + a "none" score
op per tagged number: none | diff | sum | pct | ratio | reduction (a number computed from cells)
operand per (tagged number, cell) pair: is this cell an operand of that computation?
A number is the mean of its tokens (the number and its [#k] tag); a cell is the mean of its "{...}" tokens. Pair scores
are an MLP over [m + c projections, projection of m * c]. The heads are small enough to run in numpy at inference
(decosa_api.cellpointer), so only the encoder is exported to ONNX.
"""
from __future__ import annotations
import torch
from torch import nn
OPS = ["none", "diff", "sum", "pct", "ratio", "reduction"]
class PairScorer(nn.Module):
def __init__(self, d: int, h: int = 256):
super().__init__()
self.m = nn.Linear(d, h)
self.c = nn.Linear(d, h, bias=False)
self.x = nn.Linear(d, h, bias=False)
self.out = nn.Linear(h, 1)
def forward(self, hm, hc): # (M, D), (C, D) -> (M, C)
z = self.m(hm)[:, None, :] + self.c(hc)[None, :, :] + self.x(hm[:, None, :] * hc[None, :, :])
return self.out(torch.nn.functional.gelu(z)).squeeze(-1)
class Pointer(nn.Module):
def __init__(self, encoder, d: int):
super().__init__()
self.encoder = encoder
self.drop = nn.Dropout(0.1)
self.ptr = PairScorer(d)
self.none = nn.Linear(d, 1)
self.op = nn.Linear(d, len(OPS))
self.opnd = PairScorer(d)
def encode(self, input_ids, attention_mask):
return self.encoder(input_ids=input_ids, attention_mask=attention_mask).last_hidden_state
def heads(self, H, mpool, cpool):
"""H (L, D) for one example; mpool (M, L), cpool (C, L) row-normalised pooling matrices."""
hm = self.drop(mpool @ H)
hc = self.drop(cpool @ H)
s = self.ptr(hm, hc) # (M, C)
s = torch.cat([s, self.none(hm)], dim=1) # (M, C + 1): last column is "none"
return s, self.op(hm), self.opnd(hm, hc)
def pool_matrix(spans: list[list[int]], L: int, device) -> torch.Tensor:
m = torch.zeros((len(spans), L), device=device)
for i, toks in enumerate(spans):
if toks:
m[i, toks] = 1.0 / len(toks)
return m
def token_spans(offsets, seq_ids, spans: list[tuple[int, int]], seq: int) -> list[list[int]]:
"""Character spans in text a (seq 0) or b (seq 1) -> token index lists."""
toks = [(k, a, b) for k, ((a, b), s) in enumerate(zip(offsets, seq_ids)) if s == seq and b > a]
out = []
for x, y in spans:
out.append([k for k, a, b in toks if a < y and b > x])
return out