Return an activation function given a string
(activation)
| 101 | |
| 102 | |
| 103 | def _get_activation_fn(activation): |
| 104 | """Return an activation function given a string""" |
| 105 | if activation == "relu": |
| 106 | return F.relu |
| 107 | if activation == "gelu": |
| 108 | return F.gelu |
| 109 | if activation == "glu": |
| 110 | return F.glu |
| 111 | if activation == "prelu": |
| 112 | return nn.PReLU() |
| 113 | if activation == "selu": |
| 114 | return F.selu |
| 115 | raise RuntimeError(F"activation should be relu/gelu, not {activation}.") |
| 116 | |
| 117 | |
| 118 | def _get_clones(module, N, layer_share=False): |