File size: 4,928 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
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
"""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