File size: 2,855 Bytes
2769c0b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
"""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