| 171 | |
| 172 | # Cell |
| 173 | class _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 |
| 206 | class _ExogenousBasisInterpretable(nn.Module): |