File size: 650 Bytes
76d61a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
from __future__ import annotations

import torch
from torch import nn

from ..openstl.models.simvp_model import SimVP_Model


class SimVPCI(nn.Module):
    """OpenSTL SimVP wrapper that always returns a dictionary.

    The network body is the local OpenSTL ``SimVP_Model`` copy. This wrapper only
    changes the public return type from tensor to ``{"ci": tensor}``.
    """

    def __init__(self, *args, **kwargs):
        super().__init__()
        self.backbone = SimVP_Model(*args, **kwargs)

    def forward(self, x_raw: torch.Tensor, **kwargs) -> dict[str, torch.Tensor]:
        ci = self.backbone(x_raw, **kwargs)
        return {"ci": ci}