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}
|