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

Class _IdentityBasis

models/NHits.py:159–189  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

157
158# Cell
159class _IdentityBasis(nn.Module):
160 def __init__(self, backcast_size: int, forecast_size: int, interpolation_mode: str):
161 super().__init__()
162 assert (interpolation_mode in ['linear','nearest']) or ('cubic' in interpolation_mode)
163 self.forecast_size = forecast_size
164 self.backcast_size = backcast_size
165 self.interpolation_mode = interpolation_mode
166
167 def forward(self, theta: t.Tensor, insample_x_t: t.Tensor, outsample_x_t: t.Tensor) -> Tuple[t.Tensor, t.Tensor]:
168
169 backcast = theta[:, :self.backcast_size]
170 knots = theta[:, self.backcast_size:]
171
172 if self.interpolation_mode=='nearest':
173 knots = knots[:,None,:]
174 forecast = F.interpolate(knots, size=self.forecast_size, mode=self.interpolation_mode)
175 forecast = forecast[:,0,:]
176 elif self.interpolation_mode=='linear':
177 knots = knots[:,None,:]
178 forecast = F.interpolate(knots, size=self.forecast_size, mode=self.interpolation_mode)
179 forecast = forecast[:,0,:]
180 elif 'cubic' in self.interpolation_mode:
181 batch_size = len(backcast)
182 knots = knots[:,None,None,:]
183 forecast = t.zeros((len(knots), self.forecast_size)).to(knots.device)
184 n_batches = int(np.ceil(len(knots)/batch_size))
185 for i in range(n_batches):
186 forecast_i = F.interpolate(knots[i*batch_size:(i+1)*batch_size], size=self.forecast_size, mode='bicubic')
187 forecast[i*batch_size:(i+1)*batch_size] += forecast_i[:,0,0,:]
188
189 return backcast, forecast
190
191# Cell
192class _ExogenousBasisInterpretable(nn.Module):

Callers 1

create_stackMethod · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected