Download GeometryForcing/algorithms/vae/common/modules/ops.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 660 Bytes
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/algorithms/vae/common/modules/ops.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/GeometryForcing/algorithms/vae/common/modules/ops.py
-
curl -L -o ops.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/algorithms/vae/common/modules/ops.py
660 Bytes
| from typing import Callable | |
| import torch | |
| from einops import rearrange | |
| def video_to_image(func: Callable) -> Callable: | |
| def wrapper(self, x: torch.Tensor, *args, **kwargs): | |
| if x.dim() == 5: | |
| t = x.shape[2] | |
| x = rearrange(x, "b c t h w -> (b t) c h w") | |
| x = func(self, x, *args, **kwargs) | |
| x = rearrange(x, "(b t) c h w -> b c t h w", t=t) | |
| else: | |
| x = func(self, x, *args, **kwargs) | |
| return x | |
| return wrapper | |
| def nonlinearity(x): | |
| return x * torch.sigmoid(x) | |
| def cast_tuple(t, length=1): | |
| return t if isinstance(t, tuple) or isinstance(t, list) else ((t,) * length) | |