ci-net / code /training /src /model /simvp_ci.py
lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
650 Bytes
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}