(pe, learn_pe, q_len, d_model)
| 94 | return cpe |
| 95 | |
| 96 | def positional_encoding(pe, learn_pe, q_len, d_model): |
| 97 | # Positional encoding |
| 98 | if pe == None: |
| 99 | W_pos = torch.empty((q_len, d_model)) # pe = None and learn_pe = False can be used to measure impact of pe |
| 100 | nn.init.uniform_(W_pos, -0.02, 0.02) |
| 101 | learn_pe = False |
| 102 | elif pe == 'zero': |
| 103 | W_pos = torch.empty((q_len, 1)) |
| 104 | nn.init.uniform_(W_pos, -0.02, 0.02) |
| 105 | elif pe == 'zeros': |
| 106 | W_pos = torch.empty((q_len, d_model)) |
| 107 | nn.init.uniform_(W_pos, -0.02, 0.02) |
| 108 | elif pe == 'normal' or pe == 'gauss': |
| 109 | W_pos = torch.zeros((q_len, 1)) |
| 110 | torch.nn.init.normal_(W_pos, mean=0.0, std=0.1) |
| 111 | elif pe == 'uniform': |
| 112 | W_pos = torch.zeros((q_len, 1)) |
| 113 | nn.init.uniform_(W_pos, a=0.0, b=0.1) |
| 114 | elif pe == 'lin1d': W_pos = Coord1dPosEncoding(q_len, exponential=False, normalize=True) |
| 115 | elif pe == 'exp1d': W_pos = Coord1dPosEncoding(q_len, exponential=True, normalize=True) |
| 116 | elif pe == 'lin2d': W_pos = Coord2dPosEncoding(q_len, d_model, exponential=False, normalize=True) |
| 117 | elif pe == 'exp2d': W_pos = Coord2dPosEncoding(q_len, d_model, exponential=True, normalize=True) |
| 118 | elif pe == 'sincos': W_pos = PositionalEncoding(q_len, d_model, normalize=True) |
| 119 | else: raise ValueError(f"{pe} is not a valid pe (positional encoder. Available types: 'gauss'=='normal', \ |
| 120 | 'zeros', 'zero', uniform', 'lin1d', 'exp1d', 'lin2d', 'exp2d', 'sincos', None.)") |
| 121 | return nn.Parameter(W_pos, requires_grad=learn_pe) |
no test coverage detected