Return an activation function given a string
(activation)
| 173 | |
| 174 | |
| 175 | def _get_activation_fn(activation): |
| 176 | """Return an activation function given a string""" |
| 177 | if activation == "relu": |
| 178 | return F.relu |
| 179 | if activation == "gelu": |
| 180 | return F.gelu |
| 181 | if activation == "glu": |
| 182 | return F.glu |
| 183 | raise RuntimeError(F"activation should be relu/gelu, not {activation}.") |
| 184 | |
| 185 | |
| 186 | class MLP(nn.Module): |