Return an activation function given a string.
(activation, d_model=256, batch_dim=0)
| 167 | |
| 168 | |
| 169 | def _get_activation_fn(activation, d_model=256, batch_dim=0): |
| 170 | """Return an activation function given a string.""" |
| 171 | if activation == 'relu': |
| 172 | return F.relu |
| 173 | if activation == 'gelu': |
| 174 | return F.gelu |
| 175 | if activation == 'glu': |
| 176 | return F.glu |
| 177 | if activation == 'prelu': |
| 178 | return nn.PReLU() |
| 179 | if activation == 'selu': |
| 180 | return F.selu |
| 181 | raise RuntimeError(F'activation should be relu/gelu, not {activation}.') |
| 182 | |
| 183 | |
| 184 | def gen_sineembed_for_position(pos_tensor): |