"""The joint model: encoder -> span representations -> NER head + relation head. Port of `radgraph.dygie.models.dygie.DyGIE`, coref/events removed (see package docstring in ner_head.py / relation_head.py for why: every config in this project sets their loss weight to 0, so they never trained or predicted anything in v1 either). """ from typing import Dict import torch from torch import nn from .dataset import Example from .ner_head import NERHead from .relation_head import RelationHead from .span_extractor import EndpointSpanExtractor from .tokenizer_embedder import MismatchedEmbedder from .vocab import Vocabulary def _xavier_init_weights(module: nn.Module) -> None: """xavier_normal_ on every 2-D '*.weight' / '*.weight_matrix' param -- the `module_initializer` regexes every config in configs/medbit/ applies to the NER and relation submodules (never to the pretrained encoder, which keeps its pretrained init).""" for name, param in module.named_parameters(): if param.dim() >= 2 and (name.endswith("weight") or name.endswith("weight_matrix")): nn.init.xavier_normal_(param) class DyGIEModel(nn.Module): def __init__(self, vocab: Vocabulary, encoder_name: str, max_length: int, max_span_width: int, feature_size: int, feedforward_params: dict, loss_weights: Dict[str, float], relation_spans_per_word: float, train_encoder: bool = True, span_pooling: bool = False, transformer_params: dict = None, relation_context: bool = False, relation_feedforward_params: dict = None): super().__init__() self.vocab = vocab self.loss_weights = loss_weights self.embedder = MismatchedEmbedder(encoder_name, max_length, train_encoder) self.span_extractor = EndpointSpanExtractor( self.embedder.get_output_dim(), num_width_embeddings=max_span_width, span_width_embedding_dim=feature_size, mean_pool=span_pooling) span_emb_dim = self.span_extractor.get_output_dim() self.ner = NERHead(vocab, span_emb_dim, feedforward_params, transformer_params) self.relation = RelationHead(vocab, span_emb_dim, self.embedder.get_output_dim(), feedforward_params, relation_spans_per_word, transformer_params, relation_context, relation_feedforward_params) _xavier_init_weights(self.ner) _xavier_init_weights(self.relation) nn.init.xavier_normal_(self.span_extractor.width_embedding.weight) def forward(self, example: Example) -> dict: device = next(self.embedder.parameters()).device word_embeddings = self.embedder(example.words) spans_tensor = torch.tensor(example.spans, dtype=torch.long, device=device) span_embeddings = self.span_extractor(word_embeddings, spans_tensor) zero = word_embeddings.new_zeros(()) output_ner, output_relation = {"loss": zero}, {"loss": zero} if self.loss_weights["ner"] > 0: output_ner = self.ner(example.dataset, example.spans, span_embeddings, example.sentence, example.ner_label_ids.to(device)) if self.loss_weights["relation"] > 0: output_relation = self.relation(example.dataset, example.spans, span_embeddings, len(example.words), word_embeddings, example.relation_gold, example.sentence) loss = (self.loss_weights["ner"] * output_ner.get("loss", zero) + self.loss_weights["relation"] * output_relation.get("loss", zero)) loss = loss * example.weight return {"loss": loss, "ner": output_ner, "relation": output_relation} def get_metrics(self, reset: bool = False) -> Dict[str, float]: res = {} res.update(self.ner.get_metrics(reset)) res.update(self.relation.get_metrics(reset)) return res