File size: 774 Bytes
0c48771
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
"""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)