Download code/training/src/model/simvp_ci.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 650 Bytes
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/model/simvp_ci.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/model/simvp_ci.py
-
curl -L -o simvp_ci.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/model/simvp_ci.py
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} | |