| 157 | |
| 158 | # Cell |
| 159 | class _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 |
| 192 | class _ExogenousBasisInterpretable(nn.Module): |