CLIP-Search-Edit / clip_styler.py
Tim Zhang
trying new cnn
3601849
Raw
History Blame
7.28 kB
from __future__ import annotations
from typing import Optional
from PIL import Image
import torch
import torch.optim as optim
from torchvision import transforms, models
from torchvision.transforms.functional import adjust_contrast
from transformers import CLIPModel, AutoTokenizer
from cnn1 import UNet
from util import (
DEVICE,
IMG_SIZE,
LR,
CROP_SIZE,
RESIZE,
NUM_EPOCHS,
NUM_CROPS,
PATCH_THRESHOLD,
L_TV,
L_PATCH,
L_DIR,
L_CONTENT,
normalize,
clip_normalize,
get_features,
prompt_ensemble,
get_image_prior_losses,
)
class ClipStyler:
"""
Wraps the CLIP-guided optimization loop from main.py so it can be reused
by both the CLI entrypoint and the Gradio application.
"""
def __init__(
self,
device: Optional[torch.device] = None,
num_epochs: int = NUM_EPOCHS,
) -> None:
self.device = device or DEVICE
self.num_epochs = num_epochs
# Text/image encoders
self.clip_model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32").to(
self.device
)
self.clip_model.eval()
for param in self.clip_model.parameters():
param.requires_grad_(False)
self.tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
# Perceptual backbone
self.vgg19 = models.vgg19(pretrained=True).features.to(self.device).eval()
for param in self.vgg19.parameters():
param.requires_grad_(False)
# Image helpers
self.image_transform = transforms.Compose(
[transforms.Resize((IMG_SIZE, IMG_SIZE)), transforms.ToTensor()]
)
self.to_pil = transforms.ToPILImage()
self.crop_transform = transforms.RandomCrop(CROP_SIZE)
self.augment_transform = transforms.Compose(
[
transforms.RandomPerspective(fill=0, p=1, distortion_scale=0.5),
transforms.Resize(RESIZE),
]
)
def stylize(
self,
image: Image.Image,
prompt: str,
source: str = "a Photo",
num_epochs: Optional[int] = None,
log_interval: int = 20,
) -> Image.Image:
"""
Run the CLIP-guided optimization loop and return the stylized PIL image.
"""
if image is None:
raise ValueError("Image must be provided for stylization.")
epochs = num_epochs if num_epochs is not None else self.num_epochs
base_img = self._pil_to_tensor(image)
content_features = get_features(normalize(base_img), self.vgg19)
cnn_model = UNet(text_dim=512).to(self.device)
cnn_model.train()
optimizer = optim.Adam(cnn_model.parameters(), lr=LR)
scheduler = torch.optim.lr_scheduler.StepLR(
optimizer, step_size=100, gamma=0.5
)
with torch.no_grad():
text_features = self._encode_text(prompt)
source_features = self._encode_text(source)
base_img_features = self._image_features(base_img)
style_direction = text_features - source_features
style_direction /= style_direction.norm(dim=-1, keepdim=True)
final_target = base_img
for epoch in range(epochs + 1):
scheduler.step()
target = cnn_model(base_img, text_embedding=text_features)
final_target = target
target_features = get_features(normalize(target), self.vgg19)
content_loss = torch.mean(
(target_features["conv4_2"] - content_features["conv4_2"]) ** 2
)
content_loss += torch.mean(
(target_features["conv5_2"] - content_features["conv5_2"]) ** 2
)
img_aug = self._augment_crops(target)
img_aug_features = self.clip_model.get_image_features(
pixel_values=clip_normalize(img_aug)
)
img_aug_features /= img_aug_features.norm(dim=-1, keepdim=True)
img_direction = img_aug_features - base_img_features
img_direction /= img_direction.norm(dim=-1, keepdim=True)
patch_text_direction = style_direction.repeat(img_direction.size(0), 1)
patch_text_direction /= patch_text_direction.norm(dim=-1, keepdim=True)
loss_patch = 1 - torch.cosine_similarity(
img_direction, patch_text_direction, dim=1
)
loss_patch = torch.where(
loss_patch < PATCH_THRESHOLD, torch.zeros_like(loss_patch), loss_patch
).mean()
glob_features = self.clip_model.get_image_features(
pixel_values=clip_normalize(target)
)
glob_features /= glob_features.norm(dim=-1, keepdim=True)
glob_direction = glob_features - base_img_features
glob_direction /= glob_direction.norm(dim=-1, keepdim=True)
loss_glob = (1 - torch.cosine_similarity(glob_direction, style_direction, dim=1)).mean()
loss_tv = L_TV * get_image_prior_losses(target)
total_loss = (
(L_PATCH * loss_patch)
+ (L_DIR * loss_glob)
+ (L_CONTENT * content_loss)
+ loss_tv
)
optimizer.zero_grad()
total_loss.backward()
optimizer.step()
'''
if log_interval and epoch % log_interval == 0:
print(
f"[ClipStyler] Epoch {epoch}: total={total_loss.item():.4f}, "
f"content={content_loss.item():.4f}, patch={loss_patch.item():.4f}, "
f"dir={loss_glob.item():.4f}, tv={loss_tv.item():.4f}"
)
'''
return self._tensor_to_pil(final_target)
def _pil_to_tensor(self, image: Image.Image) -> torch.Tensor:
if image.mode != "RGB":
image = image.convert("RGB")
tensor = self.image_transform(image).unsqueeze(0)
return tensor.to(self.device)
def _tensor_to_pil(self, tensor: torch.Tensor) -> Image.Image:
tensor = tensor.detach().clone().cpu().clamp(0, 1).squeeze(0)
tensor = adjust_contrast(tensor, 1.5)
return self.to_pil(tensor)
def _augment_crops(self, target: torch.Tensor) -> torch.Tensor:
crops = []
for _ in range(NUM_CROPS):
crop = self.crop_transform(target)
crop = self.augment_transform(crop)
crops.append(crop)
return torch.cat(crops, dim=0)
@torch.no_grad()
def _encode_text(self, text: str) -> torch.Tensor:
prompts = prompt_ensemble(text)
tokens = self.tokenizer(prompts, padding=True, return_tensors="pt").to(
self.device
)
text_features = self.clip_model.get_text_features(**tokens)
text_features = text_features.mean(dim=0, keepdim=True)
text_features /= text_features.norm(dim=-1, keepdim=True)
return text_features
@torch.no_grad()
def _image_features(self, tensor: torch.Tensor) -> torch.Tensor:
feats = self.clip_model.get_image_features(pixel_values=clip_normalize(tensor))
feats /= feats.norm(dim=-1, keepdim=True)
return feats