CLIP-Search-Edit / clip_styler.py
Tim Zhang
inference added
b8d244d
Raw
History Blame
2.6 kB
import torch
from PIL import Image
import torchvision.transforms as transforms
from util import *
from cnn1 import UNet
from transformers import CLIPModel, AutoTokenizer
from torchvision.transforms.functional import adjust_contrast
import torch.nn.functional as F
def load_models(weights_path):
clip_model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32").to(DEVICE)
tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
vgg19 = models.vgg19(pretrained=True).features.to(DEVICE)
for param in vgg19.parameters():
param.requires_grad_(False)
cnn_model = UNet(text_dim=512).to(DEVICE)
cnn_model.load_state_dict(torch.load(weights_path, map_location=DEVICE))
return clip_model, tokenizer, vgg19, cnn_model
def reset_cnn_model():
"""Reset CNN model weights for each new image"""
global cnn_model
if cnn_model is not None:
# Reinitialize the model
for layer in cnn_model.modules():
if isinstance(layer, (torch.nn.Conv2d, torch.nn.Linear)):
torch.nn.init.xavier_uniform_(layer.weight)
if layer.bias is not None:
torch.nn.init.zeros_(layer.bias)
def style_transfer(img, prompt, weights_path, source="a Photo"):
"""
Apply style transfer to an uploaded image
Args:
image: PIL Image uploaded by user
prompt: Text style prompt
num_steps: Number of optimization steps (fewer = faster, more = better quality)
Returns:
Stylized PIL Image
"""
if img is None:
return None
img = load_image(img).to(DEVICE)
# Load models
clip_model, tokenizer, vgg19, cnn_model = load_models(weights_path)
cnn_model.eval()
with torch.no_grad():
edited_text = prompt_ensemble(prompt)
text_tokens = tokenizer(edited_text, padding=True, return_tensors="pt").to(DEVICE)
text_features = clip_model.get_text_features(**text_tokens)
text_features = text_features.mean(axis=0, keepdim=True)
text_features /= text_features.norm(dim=-1, keepdim=True)
with torch.inference_mode():
target = cnn_model(img, text_embedding=text_features).to(DEVICE)
# Post-process output
output_image = target.clone().detach()
output_image = torch.clamp(output_image, 0, 1)
output_image = adjust_contrast(output_image, 1.5)
# Convert tensor back to PIL Image
output_image = output_image.squeeze(0).cpu()
output_image = transforms.ToPILImage()(output_image)
return output_image