escapist413/StyleFusion
2
1import argparse2import os3from datetime import datetime, timedelta4from typing import List, Iterable5 6from PIL import Image7 8import torch9from torch.utils.data import DataLoader10from torch.utils.tensorboard import SummaryWriter11from torchvision import transforms12from tqdm import tqdm13import cv214import numpy as np15 16from models import VGG, TransNet17from datasets import COCODataset18from utils import load_image, save_image, make_transform, save_model, calculate_style_loss, calculate_content_loss, \19 denormalize20 21import gradio as gr22from typing import List, Iterable23from PIL import Image24import torch25from torchvision import transforms26 27device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')28# ----------------路径参数----------------29 30# 内容特征层及loss加权系数31content_layers = {'5': 0.5, '10': 0.5} # 使用vgg的较浅层特征作为内容特征,保证生成图片内容结构相似性32# 风格特征层及loss加权系数33style_layers = {'0': 0.2, '5': 0.2, '10': 0.2, '19': 0.2, '28': 0.2} # 使用vgg不同深度的风格特征,生成风格更加层次丰富34 35transform = make_transform(size=(300, 450), normalize=True) # 图像变换36image_style = load_image('./data/udnie.jpg', transform=transform).to(device) # 风格图像37vgg = VGG(content_layers, style_layers).to(device) # 特征提取网络,只用来提取特征,不进行训练38model = TransNet(input_size=(300, 450)).to(device) # 内容生成网络,用于生成风格图片,进行训练39model.load_state_dict(torch.load('./models/udnie.pth', map_location=device))40 41 42def process_images(image) :43 image = transform(image).to(device)44 model.to(device)45 batch_generated = model(image)46 batch_generated = denormalize(batch_generated).detach().cpu()47 batch_generated = transforms.ToPILImage()(batch_generated)48 return batch_generated49 50 51# 创建 Gradio 接口52demo = gr.Interface(53 fn=process_images,54 inputs=gr.Image(type="pil", label="输入图像"),55 outputs=gr.Image(type="pil", label="生成图像")56)57 58# 启动 Gradio 界面59demo.launch(share=True)60 