Spaces:
Sleeping
Sleeping
File size: 5,746 Bytes
77b48d9 7323d20 77b48d9 7323d20 77b48d9 7323d20 77b48d9 202f334 77b48d9 202f334 77b48d9 7323d20 77b48d9 7323d20 77b48d9 f8714cf 77b48d9 202f334 77b48d9 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 | import os
import torch
from flask import Flask, render_template, request, redirect, url_for, send_from_directory
from flask_wtf import FlaskForm
from flask_bootstrap import Bootstrap
from werkzeug.utils import secure_filename
from wtforms import FileField, SubmitField, FloatField, HiddenField
from wtforms.validators import InputRequired
from PIL import Image
from torchvision import transforms
import io
#existing AdaIN code
from utils.models import VGGEncoder, Decoder
from utils.utils import adaptive_instance_normalization, calc_mean_std
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
app = Flask(__name__)
app.config['SECRET_KEY'] = 'supersecretkey'
app.config['UPLOAD_FOLDER'] = os.path.join(BASE_DIR, 'static', 'uploads')
app.config['ALLOWED_EXTENSIONS'] = {'png', 'jpg', 'jpeg'}
app.config['WTF_CSRF_ENABLED'] = False
Bootstrap(app)
app.config['UPLOAD_FOLDER'] = os.path.join(BASE_DIR, 'static', 'uploads')
os.makedirs(app.config['UPLOAD_FOLDER'], exist_ok=True)
class UploadForm(FlaskForm):
content = FileField('Content Image')
style = FileField('Style Image')
content_path = HiddenField()
style_path = HiddenField()
alpha = FloatField('Alpha', default=1.0)
submit = SubmitField('Transfer Style')
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
BASE_DIR = os.path.dirname(os.path.abspath(__file__))
encoder_path= os.path.join(BASE_DIR, "vgg_normalised.pth")
decoder_path = os.path.join(BASE_DIR, "experiment", "final_exp", "decoder_final.pth")
encoder = VGGEncoder(encoder_path).to(device)
decoder = Decoder().to(device)
decoder.load_state_dict(torch.load(decoder_path, map_location=device))
encoder.eval()
decoder.eval()
def allowed_file(filename):
return '.' in filename and \
filename.rsplit('.', 1)[1].lower() in app.config['ALLOWED_EXTENSIONS']
def style_transfer(content_image, style_image, encoder, decoder, alpha, device):
content_transform = transforms.Compose([
transforms.Resize(512),
transforms.ToTensor()
])
style_transform = transforms.Compose([
transforms.Resize(512),
transforms.ToTensor()
])
content_image = content_transform(content_image).unsqueeze(0).to(device)
style_image = style_transform(style_image).unsqueeze(0).to(device)
with torch.no_grad():
content_feats = encoder(content_image, is_test=True)
style_feats = encoder(style_image, is_test=True)
stylized_feats = adaptive_instance_normalization(content_feats, style_feats)
stylized_feats = alpha * stylized_feats + (1 - alpha) * content_feats
stylized_image = decoder(stylized_feats)
return stylized_image
def save_image(image, path):
image = image.cpu().clone()
image = image.squeeze(0)
image = image.clamp(0, 1)
image = transforms.ToPILImage()(image)
image.save(path)
@app.route('/', methods=['GET', 'POST'])
def index():
form = UploadForm()
result_image = None
content_filename = None
style_filename = None
error = None
if form.validate_on_submit():
if form.content.data and form.content.data.filename:
if allowed_file(form.content.data.filename):
content_filename = secure_filename(form.content.data.filename)
form.content.data.save(os.path.join(app.config['UPLOAD_FOLDER'], content_filename))
form.content_path.data = content_filename
else:
content_filename = form.content_path.data
if form.style.data and form.style.data.filename:
if allowed_file(form.style.data.filename):
style_filename = secure_filename(form.style.data.filename)
form.style.data.save(os.path.join(app.config['UPLOAD_FOLDER'], style_filename))
form.style_path.data = style_filename
else:
style_filename = form.style_path.data
if content_filename and style_filename:
content_path = os.path.join(app.config['UPLOAD_FOLDER'], content_filename)
style_path = os.path.join(app.config['UPLOAD_FOLDER'], style_filename)
try:
content_image = Image.open(content_path).convert('RGB')
style_image = Image.open(style_path).convert('RGB')
alpha = float(form.alpha.data)
stylized_image = style_transfer(content_image, style_image, encoder, decoder, alpha, device)
result_filename = 'stylized_' + content_filename
result_path = os.path.join(app.config['UPLOAD_FOLDER'], result_filename)
save_image(stylized_image, result_path)
result_image = result_filename
except Exception as e:
import traceback
error = traceback.format_exc()
print(error)
else:
if not content_filename:
error = 'Please upload content image'
if not style_filename:
error = 'Please upload style image'
return render_template('index.html', form=form, result_image=result_image, content_image=content_filename,
style_image=style_filename, error=error)
@app.route('/uploads/<filename>')
def send_image(filename):
full_path = os.path.join(app.config['UPLOAD_FOLDER'], filename)
print(f"Serving file: {filename}, exists: {os.path.exists(full_path)}")
return send_from_directory(app.config['UPLOAD_FOLDER'], filename)
@app.route('/examples/<path:filename>')
def send_example(filename):
examples_dir = os.path.join(BASE_DIR, 'examples')
return send_from_directory('examples', filename)
if __name__ == '__main__':
app.run(host="0.0.0.0", port=7860)
|