Ermolaev / model.py
tuva's picture
Publish Ermolaev CUT model, documentation, examples, and inference code
22899f3 verified
Raw History Blame Contribute Delete
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()