Radgraph-IT / pruner.py
roccoangelella's picture
Add transformers-compatible wrapper (AutoModel/AutoTokenizer via trust_remote_code)
0c48771 verified
Raw
History Blame Contribute Delete
1.08 kB
"""Mention pruner: score every span, keep the top-k, in original span order.
Port of `radgraph.dygie.models.entity_beam_pruner.Pruner`'s default path (no entity_beam, no
gold_beam -- this project's relation module never sets either, see relation_head.py).
Simplified for batch_size=1: no masking, no batch dimension.
"""
from typing import Tuple
import torch
from torch import nn
class Pruner(nn.Module):
def __init__(self, scorer: nn.Module):
super().__init__()
self.scorer = scorer # (num_spans, dim) -> (num_spans, 1)
def forward(self, span_embeddings: torch.Tensor, num_items_to_keep: int
) -> Tuple[torch.Tensor, torch.LongTensor, torch.Tensor]:
scores = self.scorer(span_embeddings).squeeze(-1) # (num_spans,)
k = max(1, min(num_items_to_keep, span_embeddings.size(0)))
_, top_indices = scores.topk(k)
top_indices, _ = torch.sort(top_indices)
top_scores = scores[top_indices]
top_embeddings = span_embeddings[top_indices]
return top_embeddings, top_indices, top_scores