decosa-cell-pointer-modernbert-base / cellpointer_input.py
sipratt's picture
decosa-cell-pointer-modernbert-base: weights, card, eval summary (Apache-2.0)
2769c0b verified
Raw History Blame Contribute Delete
4.93 kB
"""The cell-pointer's input: sentences with their numbers tagged, and candidate tables written out cell by cell, with the
character span of every tagged number and every cell. Shared by training (scripts/cell_pointer/) and the runtime
(decosa_api.cellpointer), so both see exactly the same text.
text_a "11.4.1 Primary Efficacy Endpoint: PASI 75 was achieved by 128 [#1] of 205 [#2] (62.4% [#3]) ..."
text_b "<T5> Table 14.2.1: Primary endpoint: PASI 75 at Week 16 (full analysis set)
~ (column header N): {Placebo} {Zenavotide 40 mg} {Zenavotide 80 mg}
~ Responders, n (%): {Placebo} {Zenavotide 40 mg} {Zenavotide 80 mg}
..."
Value-blind (the default): a cell is written as its column label only, and "(N=206)" is dropped from column labels, so the
model can only choose a cell by what the sentence says (group, row, timepoint, population), never by the number printed
in it. Value-aware (blind=False) writes "{Placebo = 52 (25.2)}". Only cells that hold a number are candidates.
"""
from __future__ import annotations
import re
from dataclasses import dataclass, field
import csr_tables as TB
N_IN_LABEL = re.compile(r"\(?\s*[Nn]\s*=\s*[\d,]+\s*\)?")
MAX_LABEL = 80
@dataclass
class Rendered:
text_a: str
text_b: str
mentions: list[tuple[str, int, int]] # (tag "#1", start, end) in text_a: the number and its tag
cells: dict[tuple[str, str], tuple[int, int]] # (table key, cell ref) -> (start, end) in text_b
tables: list[str] = field(default_factory=list) # table keys, in order
def col_label(c: TB.Cell, blind: bool) -> str:
lab = (c.arm or c.col or "").strip()
if blind:
lab = N_IN_LABEL.sub("", lab)
lab = re.sub(r"\s+", " ", lab).strip(" ,;:/-") or f"column {c.c}"
return lab[:MAX_LABEL]
def row_label(c: TB.Cell) -> str:
if c.row == "(column header)":
return "(column header N)"
return (re.sub(r"\s+", " ", c.row or "").strip() or f"row {c.r}")[:120]
def render_table(t: TB.Table, blind: bool = True, max_rows: int | None = None) -> tuple[str, dict[str, tuple[int, int]]]:
"""One table: the caption line, then one line per row that holds numbers. Returns (text, {ref: (start, end)})."""
head = f"<{t.key}> {t.label()}: {t.title}" + (f" (copy of Table {t.source})" if t.source else "")
out = [head]
pos = len(head)
spans: dict[str, tuple[int, int]] = {}
rows: dict[int, list[TB.Cell]] = {}
for c in t.cells.values():
if c.numbers:
rows.setdefault(c.r, []).append(c)
for k, r in enumerate(sorted(rows)):
if max_rows is not None and k >= max_rows:
break
cells = sorted(rows[r], key=lambda c: c.c)
line = "\n~ " + row_label(cells[0]) + ":"
for c in cells:
lab = col_label(c, blind)
piece = "{" + lab + ("" if blind else " = " + re.sub(r"\s+", " ", c.text)[:60]) + "}"
a = pos + len(line) + 1
line += " " + piece
spans[c.ref] = (a, a + len(piece))
out.append(line)
pos += len(line)
return "".join(out), spans
def render(sents: list[dict], tables: list[TB.Table], *, blind: bool = True, context: str = "",
max_rows: int | None = None) -> Rendered:
"""sents: [{"text", "mentions": [{"start", "end", "checked", ...}]}] (mentions.mentions per sentence). Each checked
mention gets a tag #1..#n in order. context: a heading to put first ("11.4.1 Primary Efficacy Endpoint")."""
a = (context.strip() + ": ") if context.strip() else ""
mentions: list[tuple[str, int, int]] = []
k = 0
for s in sents:
if a and not a.endswith(" "):
a += " "
last = 0
txt = s["text"]
for m in s["mentions"]:
if not m.get("checked"):
continue
k += 1
a += txt[last:m["start"]]
st = len(a)
a += txt[m["start"]:m["end"]] + f" [#{k}]"
mentions.append((f"#{k}", st, len(a)))
last = m["end"]
a += txt[last:]
b = ""
cells: dict[tuple[str, str], tuple[int, int]] = {}
for t in tables:
if b:
b += "\n"
txt, spans = render_table(t, blind=blind, max_rows=max_rows)
off = len(b)
b += txt
for ref, (x, y) in spans.items():
cells[(t.key, ref)] = (off + x, off + y)
return Rendered(a, b, mentions, cells, [t.key for t in tables])
def table_from_rows(key: str, tid: str, title: str, rows: list[list], header_rows: int | None = None, source: str | None = None) -> TB.Table:
"""A plain grid (first column: row labels) as a table the pointer can read; for callers outside the CSR verifier
(the number-consistency checker, filing tie-out)."""
t = TB.Table.from_rows(tid, title, rows, header_rows=header_rows, kind="tlf", source=source)
t.key = key
return t