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}