Download model.py from tuva/Ermolaev: direct link, hf CLI and curl.
- Browser
- Download file 2.1 kB
-
https://huggingface.co/tuva/Ermolaev/resolve/main/model.py
- Command line
-
hf download hf://tuva/Ermolaev/model.py
-
curl -L -o model.py https://huggingface.co/tuva/Ermolaev/resolve/main/model.py
2.1 kB
| """Standalone PyTorch definition for the Ermolaev CUT generator. | |
| The generator structure is derived from the official CUT implementation: | |
| https://github.com/taesungp/contrastive-unpaired-translation | |
| """ | |
| from __future__ import annotations | |
| import functools | |
| import torch | |
| from torch import nn | |
| class ResnetBlock(nn.Module): | |
| def __init__(self, dim: int): | |
| super().__init__() | |
| norm = functools.partial(nn.InstanceNorm2d, affine=False, track_running_stats=False) | |
| self.conv_block = nn.Sequential( | |
| nn.ReflectionPad2d(1), nn.Conv2d(dim, dim, 3, bias=True), norm(dim), nn.ReLU(True), | |
| nn.ReflectionPad2d(1), nn.Conv2d(dim, dim, 3, bias=True), norm(dim), | |
| ) | |
| def forward(self, x): | |
| return x + self.conv_block(x) | |
| class ErmolaevGenerator(nn.Module): | |
| """3-channel, 6-block CUT ResNet generator with 32 base filters.""" | |
| def __init__(self): | |
| super().__init__() | |
| norm = functools.partial(nn.InstanceNorm2d, affine=False, track_running_stats=False) | |
| layers = [ | |
| nn.ReflectionPad2d(3), nn.Conv2d(3, 32, 7, bias=True), norm(32), nn.ReLU(True), | |
| nn.Conv2d(32, 64, 3, stride=2, padding=1, bias=True), norm(64), nn.ReLU(True), | |
| nn.Conv2d(64, 128, 3, stride=2, padding=1, bias=True), norm(128), nn.ReLU(True), | |
| ] | |
| layers.extend(ResnetBlock(128) for _ in range(6)) | |
| layers.extend([ | |
| nn.ConvTranspose2d(128, 64, 3, stride=2, padding=1, output_padding=1, bias=True), | |
| norm(64), nn.ReLU(True), | |
| nn.ConvTranspose2d(64, 32, 3, stride=2, padding=1, output_padding=1, bias=True), | |
| norm(32), nn.ReLU(True), | |
| nn.ReflectionPad2d(3), nn.Conv2d(32, 3, 7), nn.Tanh(), | |
| ]) | |
| self.model = nn.Sequential(*layers) | |
| def forward(self, image): | |
| return self.model(image) | |
| def load_generator(weights_path, device="cpu"): | |
| model = ErmolaevGenerator() | |
| state = torch.load(weights_path, map_location="cpu", weights_only=True) | |
| model.load_state_dict(state, strict=True) | |
| return model.to(device).eval() | |