File size: 2,793 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
"""Minimal PyTorch usage of the Decosa cell-pointer (CPU is fine).

    pip install torch transformers safetensors
    python usage.py                     # the value-blind model (the default)
    MODEL_DIR=aware python usage.py     # the value-aware twin, for comparison

It renders one sentence and one table exactly as the model was trained (cellpointer_input.py), runs the encoder and the
heads, and prints the cell each tagged number reports. Your own code then compares the written number with that cell's
value; the model never sees the values (value-blind view).
"""
import json, os, sys
import torch
from safetensors.torch import load_file
from transformers import AutoConfig, AutoModel, AutoTokenizer

here = os.path.dirname(os.path.abspath(__file__))
d = os.path.join(here, os.environ.get("MODEL_DIR", "."))
sys.path.insert(0, here)
import cellpointer_input as CI
from modeling_cellpointer import Pointer, OPS, pool_matrix, token_spans

meta = json.load(open(os.path.join(d, "pointer.json")))
blind = meta.get("view", "blind") == "blind"
cfg = AutoConfig.from_pretrained(d)
cfg._attn_implementation = "eager"
model = Pointer(AutoModel.from_config(cfg), cfg.hidden_size)
model.load_state_dict(load_file(os.path.join(d, "model.safetensors")))
model.eval()
tok = AutoTokenizer.from_pretrained(d)

table = CI.table_from_rows("T1", "14.2.1", "Primary endpoint: PASI 75 at Week 16 (full analysis set)", [
    ["", "Placebo (N=103)", "Drug X 40 mg (N=102)"],
    ["Responders, n (%)", "18 (17.5)", "64 (62.7)"],
    ["Difference vs placebo, % (95% CI)", "", "45.2 (33.1, 57.3)"],
], header_rows=1)
text = "PASI 75 at Week 16 was reached by 64 of 102 patients (62.7%) on Drug X 40 mg and 18 (17.5%) on placebo."
mentions = []
for n in ["64", "62.7", "18", "17.5"]:
    i = text.index(n)
    mentions.append({"start": i, "end": i + len(n), "checked": True})
r = CI.render([{"text": text, "mentions": mentions}], [table], blind=blind)

enc = tok(r.text_a, r.text_b, return_offsets_mapping=True, truncation=True, max_length=meta["max_len"], return_tensors="pt")
seq = enc.sequence_ids(0)
offs = enc["offset_mapping"][0].tolist()
cells = list(r.cells)
with torch.no_grad():
    H = model.encode(enc["input_ids"], enc["attention_mask"])[0]
    L = H.shape[0]
    mp = pool_matrix(token_spans(offs, seq, [(a, b) for _, a, b in r.mentions], 0), L, H.device)
    cp = pool_matrix(token_spans(offs, seq, [r.cells[c] for c in cells], 1), L, H.device)
    s, op, _ = model.heads(H, mp, cp)
    p = torch.softmax(s / meta.get("temperature", 1.0), dim=1)
for k, (tag, a, b) in enumerate(r.mentions):
    j = int(p[k].argmax())
    where = "no cell" if j == len(cells) else f"{cells[j][0]} {cells[j][1]}"
    print(r.text_a[a:b], "->", where, f"p={float(p[k, j]):.2f}", "op=" + OPS[int(op[k].argmax())])