"""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