Download GeometryForcing/utils/torch_utils.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 903 Bytes
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/utils/torch_utils.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/GeometryForcing/utils/torch_utils.py
-
curl -L -o torch_utils.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/utils/torch_utils.py
903 Bytes
| """ | |
| This repo is forked from [Boyuan Chen](https://boyuan.space/)'s research | |
| template [repo](https://github.com/buoyancy99/research-template). | |
| By its MIT license, you must keep the above sentence in `README.md` | |
| and the `LICENSE` file to credit the author. | |
| """ | |
| from typing import Optional | |
| import torch | |
| from torch.types import _size | |
| import torch.nn as nn | |
| def freeze_model(model: nn.Module) -> None: | |
| """Freeze the torch model""" | |
| model.eval() | |
| for param in model.parameters(): | |
| param.requires_grad = False | |
| def bernoulli_tensor( | |
| size: _size, | |
| p: float, | |
| device: Optional[torch.device] = None, | |
| generator: Optional[torch.Generator] = None, | |
| ): | |
| """ | |
| Generate a tensor of the given size, | |
| where each element is sampled from a Bernoulli distribution with probability `p`. | |
| """ | |
| return torch.bernoulli(torch.full(size, p, device=device), generator=generator) | |