"""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