"""Small MLP used by every scorer (NER, mention pruner, relation): Linear -> ReLU -> Dropout, one block per entry in `hidden_dims`, each block's width taken from that entry (see configs/medbit_span_pooling/fold0.json: feedforward_params = {hidden_dims: [300, 150], dropout: 0.4}). """ from typing import List from torch import nn class FeedForward(nn.Module): def __init__(self, input_dim: int, hidden_dims: List[int], dropout: float): super().__init__() layers = [] prev = input_dim for dim in hidden_dims: layers += [nn.Linear(prev, dim), nn.ReLU(), nn.Dropout(dropout)] prev = dim self.net = nn.Sequential(*layers) self.output_dim = prev def forward(self, x): return self.net(x)