sipratt's picture
decosa-cell-pointer-modernbert-base: weights, card, eval summary (Apache-2.0)
2769c0b verified
Raw History Blame Contribute Delete
2.79 kB
"""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())])