"""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()