Return an activation function given a string.
(activation)
| 704 | |
| 705 | |
| 706 | def _get_activation_fn(activation): |
| 707 | """Return an activation function given a string.""" |
| 708 | if activation == 'relu': |
| 709 | return F.relu |
| 710 | if activation == 'gelu': |
| 711 | return F.gelu |
| 712 | if activation == 'glu': |
| 713 | return F.glu |
| 714 | if activation == 'leaky_relu': |
| 715 | return F.leaky_relu |
| 716 | raise RuntimeError(F'activation should be relu/gelu, not {activation}.') |