MCPcopy Create free account
hub / github.com/BorealisAI/scaleformer / _IdentityBasis

Class _IdentityBasis

models/NHitsMS.py:173–203  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

171
172# Cell
173class _IdentityBasis(nn.Module):
174 def __init__(self, backcast_size: int, forecast_size: int, interpolation_mode: str):
175 super().__init__()
176 assert (interpolation_mode in ['linear','nearest']) or ('cubic' in interpolation_mode)
177 self.forecast_size = forecast_size
178 self.backcast_size = backcast_size
179 self.interpolation_mode = interpolation_mode
180
181 def forward(self, theta: t.Tensor, insample_x_t: t.Tensor, outsample_x_t: t.Tensor) -> Tuple[t.Tensor, t.Tensor]:
182
183 backcast = theta[:, :self.backcast_size]
184 knots = theta[:, self.backcast_size:]
185
186 if self.interpolation_mode=='nearest':
187 knots = knots[:,None,:]
188 forecast = F.interpolate(knots, size=self.forecast_size, mode=self.interpolation_mode)
189 forecast = forecast[:,0,:]
190 elif self.interpolation_mode=='linear':
191 knots = knots[:,None,:]
192 forecast = F.interpolate(knots, size=self.forecast_size, mode=self.interpolation_mode)
193 forecast = forecast[:,0,:]
194 elif 'cubic' in self.interpolation_mode:
195 batch_size = len(backcast)
196 knots = knots[:,None,None,:]
197 forecast = t.zeros((len(knots), self.forecast_size)).to(knots.device)
198 n_batches = int(np.ceil(len(knots)/batch_size))
199 for i in range(n_batches):
200 forecast_i = F.interpolate(knots[i*batch_size:(i+1)*batch_size], size=self.forecast_size, mode='bicubic')
201 forecast[i*batch_size:(i+1)*batch_size] += forecast_i[:,0,0,:]
202
203 return backcast, forecast
204
205# Cell
206class _ExogenousBasisInterpretable(nn.Module):

Callers 1

create_stackMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected