python
import torch
from torchvision import transforms
from PIL import Image
import torchvision.utils as utils
from models import UnetGenerator 假设这是定义好的Unet生成器
加载预训练模型
def load_model(model_path):
model = UnetGenerator(3, 1, 8, 64, norm_layer=torch.nn.BatchNorm2d, use_dropout=False)
model.load_state_dict(torch.load(model_path))
return model