File size: 4,617 Bytes
0c48771
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""transformers-compatible wrapper around `DyGIEModel` (training_v2/src/dygie/model.py).

`RadgraphModel.model` is a plain `DyGIEModel`, so its state_dict keys are `model.<...>`,
matching the `model.` prefix applied when this repo's checkpoint was converted from the
original `best.pt` (`{"epoch": ..., "model_state_dict": ...}`, see observability.py's
`HFCheckpointUploader` / trainer.py in the training repo) to `model.safetensors`.

Two ways to run inference:
  - `model.predict(text)` -- raw report text in, decoded entities/relations out. Applies the
    exact word-level preprocessing training used (regex cleanup + wordpunct tokenization,
    ported from radgraph/radgraph/utils.py's `radgraph_xl_preprocess_report`, verified against
    the gold tokenization by check_tokenization_raw_it.py in the training repo).
  - `model.predict_words(words)` -- skip preprocessing if you already have the exact
    word-tokenized list (e.g. re-running on the original training data).
"""
import re
from typing import List

import torch
from transformers import PreTrainedModel
from nltk.tokenize import wordpunct_tokenize

from .configuration_radgraph import RadgraphConfig
from .document import Document
from .model import DyGIEModel
from .predictor import predict_document
from .vocab import Vocabulary


def _preprocess_report(text: str) -> str:
    """Exact port of radgraph/radgraph/utils.py:34-48 (inference-time preprocessing)."""
    text = text.replace("\\n", "  ")
    text = text.replace("\\f", "  ")
    text = text.replace("\\u2122", "      ")
    text = text.replace("\n", " ")
    text = text.replace("\\\"", "``")
    text_sub = re.sub(r"\s+", " ", text)
    t = " ".join(wordpunct_tokenize(text_sub))
    t = t.replace(").", ") .")
    t = t.replace("%.", "% .")
    t = t.replace(".'", ". '")
    t = t.replace("%,", "% ,")
    t = t.replace("%)", "% )")
    return t


def _build_vocab(config: RadgraphConfig) -> Vocabulary:
    vocab = Vocabulary()
    ner_ns = f"{config.dataset}__ner_labels"
    rel_ns = f"{config.dataset}__relation_labels"
    vocab.add_namespace(ner_ns)
    vocab.add_namespace(rel_ns)
    for token, _ in sorted(config.ner_labels.items(), key=lambda kv: kv[1]):
        if token != "":
            vocab.add_token(token, ner_ns)
    for token, _ in sorted(config.relation_labels.items(), key=lambda kv: kv[1]):
        if token != "":
            vocab.add_token(token, rel_ns)
    return vocab


class RadgraphModel(PreTrainedModel):
    config_class = RadgraphConfig

    def __init__(self, config: RadgraphConfig):
        super().__init__(config)
        self.vocab = _build_vocab(config)
        self.model = DyGIEModel(
            vocab=self.vocab,
            encoder_name=config.encoder_name,
            max_length=config.max_length,
            max_span_width=config.max_span_width,
            feature_size=config.feature_size,
            feedforward_params=config.feedforward_params,
            loss_weights=config.loss_weights,
            relation_spans_per_word=config.relation_spans_per_word,
            train_encoder=config.train_encoder,
            span_pooling=config.span_pooling,
            transformer_params=config.transformer_params,
            relation_context=config.relation_context,
            relation_feedforward_params=config.relation_feedforward_params,
        )
        # Required by transformers' from_pretrained machinery (sets up all_tied_weights_keys
        # etc.) -- harmless here since DyGIEModel has no tied weights and the checkpoint's
        # state dict gets loaded over whatever this touches anyway.
        self.post_init()

    def forward(self, words: List[str], doc_key: str = "doc"):
        """Low-level forward pass: pre-tokenized words in, raw model output dict out
        (see DyGIEModel.forward / predict_document)."""
        doc = Document.from_json({"doc_key": doc_key, "dataset": self.config.dataset, "sentences": [words]})
        return predict_document(self.model, self.vocab, doc, self.config.max_span_width)

    @torch.no_grad()
    def predict_words(self, words: List[str], doc_key: str = "doc") -> dict:
        result = self.forward(words, doc_key=doc_key)
        return {
            "doc_key": result["doc_key"],
            "words": words,
            "entities": [e.to_json() for e in result["predicted_ner"]],
            "relations": [r.to_json() for r in result["predicted_relations"]],
        }

    @torch.no_grad()
    def predict(self, text: str, doc_key: str = "doc") -> dict:
        words = _preprocess_report(text).split()
        return self.predict_words(words, doc_key=doc_key)