"""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())])