MCPcopy Create free account
hub / github.com/Anoise/WTFlib / positional_encoding

Function positional_encoding

LDPS_Graph/layers/PatchTST_layers.py:96–121  ·  view source on GitHub ↗
(pe, learn_pe, q_len, d_model)

Source from the content-addressed store, hash-verified

94 return cpe
95
96def 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)

Callers 1

__init__Method · 0.85

Calls 3

Coord1dPosEncodingFunction · 0.85
Coord2dPosEncodingFunction · 0.85
PositionalEncodingFunction · 0.70

Tested by

no test coverage detected