Radgraph-IT / ner_head.py
roccoangelella's picture
Add transformers-compatible wrapper (AutoModel/AutoTokenizer via trust_remote_code)
0c48771 verified
Raw
History Blame Contribute Delete
5 kB
"""NER head: score every candidate span, cross-entropy against the gold label (or null).
Port of `radgraph.dygie.models.ner.NERTagger`. Simplified for batch_size=1 (dataset.py): no
span mask, no per-sentence looping -- one document's worth of spans at a time.
`transformer_params` (medbit_transformer arm, see notes/v2_plan.md item 1) swaps the plain
FeedForward scorer for a TransformerBlock: every candidate span attends to every other
candidate span in the document before classification, instead of each span being scored in
isolation. NER scores every candidate span unpruned, so this is full O(n^2) self-attention
over the whole doc's span pool (accepted cost, not capped -- see v2_plan.md item 1).
"""
from typing import Dict, List
import torch
from torch import nn
from .document import PredictedNER, Sentence
from .feedforward import FeedForward
from .metrics import NERMetrics
from .transformer_block import TransformerBlock
from .vocab import Vocabulary
class NERHead(nn.Module):
def __init__(self, vocab: Vocabulary, span_emb_dim: int, feedforward_params: dict,
transformer_params: dict = None):
super().__init__()
self.vocab = vocab
self.namespaces = [ns for ns in vocab.get_namespaces() if ns.endswith("ner_labels")]
self.n_labels = {ns: vocab.get_vocab_size(ns) for ns in self.namespaces}
self.use_transformer = transformer_params is not None
self.scorers = nn.ModuleDict()
self.classifiers = nn.ModuleDict()
self.metrics: Dict[str, NERMetrics] = {}
for ns in self.namespaces:
if self.use_transformer:
body = TransformerBlock(span_emb_dim, **transformer_params)
else:
body = FeedForward(span_emb_dim, feedforward_params["hidden_dims"],
feedforward_params["dropout"])
self.scorers[ns] = body
self.classifiers[ns] = nn.Linear(body.output_dim, self.n_labels[ns] - 1)
self.metrics[ns] = NERMetrics(self.n_labels[ns])
self._loss = nn.CrossEntropyLoss(reduction="sum")
def forward(self, dataset: str, spans: List[tuple], span_embeddings: torch.Tensor,
sentence: Sentence, ner_label_ids: torch.Tensor = None) -> dict:
ns = f"{dataset}__ner_labels"
if self.use_transformer:
starts = torch.tensor([s[0] for s in spans], dtype=torch.long,
device=span_embeddings.device)
hidden = self.scorers[ns](span_embeddings, starts)
else:
hidden = self.scorers[ns](span_embeddings)
scores = self.classifiers[ns](hidden) # (num_spans, n_labels - 1)
dummy = scores.new_zeros(scores.size(0), 1) # null-label score, fixed at 0
scores = torch.cat([dummy, scores], dim=-1) # (num_spans, n_labels)
predicted = scores.argmax(dim=-1)
predictions = self.decode(ns, scores.detach(), spans, sentence)
output = {"namespace": ns, "scores": scores, "predicted": predicted, "predictions": predictions}
if ner_label_ids is not None:
self.metrics[ns](predicted, ner_label_ids)
output["loss"] = self._loss(scores, ner_label_ids)
return output
def decode(self, namespace: str, scores: torch.Tensor, spans: List[tuple],
sentence: Sentence) -> List[PredictedNER]:
softmax_scores = torch.softmax(scores, dim=-1)
raw, predicted = scores.max(dim=-1)
soft, _ = softmax_scores.max(dim=-1)
predictions = []
for i in (predicted != 0).nonzero(as_tuple=True)[0].tolist():
label = self.vocab.get_token_from_index(int(predicted[i]), namespace)
start, end = spans[i]
entry = [start, end, label, float(raw[i]), float(soft[i])]
predictions.append(PredictedNER(entry, sentence, sentence_offsets=True))
return predictions
def get_metrics(self, reset: bool = False) -> Dict[str, float]:
res = {}
for ns, metric in self.metrics.items():
prefix = ns.replace("_labels", "")
per_class = metric.get_per_class() # {class idx -> F1}, read BEFORE the reset below
for idx, class_f1 in per_class.items():
res[f"ner_perclass/{self.vocab.get_token_from_index(idx, ns)}"] = class_f1
res["ner_macro_f1"] = sum(per_class.values()) / len(per_class) if per_class else 0.0
precision, recall, f1 = metric.get_metric(reset)
res[f"{prefix}_precision"] = precision
res[f"{prefix}_recall"] = recall
res[f"{prefix}_f1"] = f1
for name in ("precision", "recall", "f1"):
values = [v for k, v in res.items()
if k.endswith(f"_{name}") and "perclass/" not in k and "_macro_" not in k]
res[f"MEAN__ner_{name}"] = sum(values) / len(values) if values else 0.0
return res