I want to replace the image editing part of InstructN2N, which uses InstructPix2Pix, with another model, CLIPStyler, that requires 200 iterations for training. My CLIPStyler works fine when run independently, but after replacing InstructPix2Pix, the output of CLIPStyler is the same for every iteration, even though the gradients are calculated successfully and the model parameters are changing. I don’t understand why the output is the same every time. This is preventing me from successfully editing the image in InstructN2N. Here is my CLIPStyler code used to replace InstructPix2Pix in the source files.
self.lambda_tv = 2e-3
self.lambda_patch = 9000
self.lambda_dir = 500
self.lambda_c = 150
self.crop_size = 128
self.content_weight = self.lambda_c
self.num_crops = 64
self.show_every = 100
self.img_width = 512
self.img_height = 512
self.steps = 200
self.lr = 5#e-4
self.thresh = 0.7
self.device = device
self.text = text
assert (self.img_width % 8) == 0, "width must be multiple of 8"
assert (self.img_height % 8) == 0, "height must be multiple of 8"
# Load pretrained VGG19 model
#self.VGG = models.vgg19(pretrained=True).features.to(self.device)
self.VGG = models.vgg19(weights=VGG19_Weights.DEFAULT).features.to(self.device)
for parameter in self.VGG.parameters():
parameter.requires_grad_(False)
# Initialize StyleNet model
self.style_net = in2n.StyleNet.UNet().to(self.device)
self.optimizer = optim.Adam(self.style_net.parameters(), lr=self.lr)
self.scheduler = torch.optim.lr_scheduler.StepLR(self.optimizer, step_size=100, gamma=0.5)
# CLIP model initialization
self.clip_model, self.preprocess = clip.load('ViT-B/32', self.device, jit=False)
# Other initialization
self.content_loss_epoch = []
self.style_loss_epoch = []
self.total_loss_epoch = []
def img_normalize(self, image):
mean = torch.tensor([0.485, 0.456, 0.406]).to(self.device).view(1, -1, 1, 1)
std = torch.tensor([0.229, 0.224, 0.225]).to(self.device).view(1, -1, 1, 1)
return (image - mean) / std
def clip_normalize(self, image):
image = F.interpolate(image, size=224, mode='bicubic')
mean = torch.tensor([0.48145466, 0.4578275, 0.40821073]).to(self.device).view(1, -1, 1, 1)
std = torch.tensor([0.26862954, 0.26130258, 0.27577711]).to(self.device).view(1, -1, 1, 1)
return (image - mean) / std
def get_image_prior_losses(self, inputs_jit):
diff1 = inputs_jit[:, :, :, :-1] - inputs_jit[:, :, :, 1:]
diff2 = inputs_jit[:, :, :-1, :] - inputs_jit[:, :, 1:, :]
diff3 = inputs_jit[:, :, 1:, :-1] - inputs_jit[:, :, :-1, 1:]
diff4 = inputs_jit[:, :, :-1, :-1] - inputs_jit[:, :, 1:, 1:]
return torch.norm(diff1) + torch.norm(diff2) + torch.norm(diff3) + torch.norm(diff4)
def compose_text_with_templates(self, text, templates):
return [template.format(text) for template in templates]
def load_content_image(self, c_path = "test_set/farm1.jpg"):
content_path = c_path
content_image = utils.load_image2(content_path, img_height=self.img_height, img_width=self.img_width)
content_image = content_image.to(self.device)
return content_image
def get_features(self, image):
return utils.get_features(self.img_normalize(image), self.VGG)
def edit_image(self, content_image):
# content_image: torch.Size([1, 3, 512, 512])<class 'torch.Tensor'>
content_features = self.get_features(content_image)
target = content_image.clone().requires_grad_(True).to(self.device)
# Initialize text and image features for CLIP
prompt = self.text
source = "a Photo"
with torch.no_grad():
template_text = self.compose_text_with_templates(prompt, imagenet_templates)
tokens = clip.tokenize(template_text).to(self.device)
text_features = self.clip_model.encode_text(tokens).detach()
text_features = text_features.mean(axis=0, keepdim=True)
text_features /= text_features.norm(dim=-1, keepdim=True)
template_source = self.compose_text_with_templates(source, imagenet_templates)
tokens_source = clip.tokenize(template_source).to(self.device)
text_source = self.clip_model.encode_text(tokens_source).detach()
text_source = text_source.mean(axis=0, keepdim=True)
text_source /= text_source.norm(dim=-1, keepdim=True)
source_features = self.clip_model.encode_image(self.clip_normalize(content_image)).detach()
source_features /= source_features.norm(dim=-1, keepdim=True)
#target = self.style_net(content_image, use_sigmoid=True).to(self.device)
for epoch in range(self.steps + 1):
self.scheduler.step()
target = self.style_net(content_image, use_sigmoid=True).to(self.device)
target.requires_grad_(True)
target_features = self.get_features(target)
content_loss = 0
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)
loss_patch = 0
img_proc = []
for n in range(self.num_crops):
target_crop = transforms.RandomCrop(self.crop_size)(target)
target_crop = transforms.RandomPerspective(fill=0, p=1, distortion_scale=0.5)(target_crop)
target_crop = transforms.Resize(224)(target_crop)
img_proc.append(target_crop)
img_proc = torch.cat(img_proc, dim=0)
img_aug = img_proc
image_features = self.clip_model.encode_image(self.clip_normalize(img_aug))
image_features /= image_features.clone().norm(dim=-1, keepdim=True)
img_direction = (image_features - source_features)
img_direction /= img_direction.clone().norm(dim=-1, keepdim=True)
text_direction = (text_features - text_source).repeat(image_features.size(0), 1)
text_direction /= text_direction.norm(dim=-1, keepdim=True)
loss_temp = (1 - torch.cosine_similarity(img_direction, text_direction, dim=1))
loss_temp[loss_temp < self.thresh] = 0
loss_patch += loss_temp.mean()
glob_features = self.clip_model.encode_image(self.clip_normalize(target))
glob_features /= glob_features.clone().norm(dim=-1, keepdim=True)
glob_direction = (glob_features - source_features)
glob_direction /= glob_direction.clone().norm(dim=-1, keepdim=True)
loss_glob = (1 - torch.cosine_similarity(glob_direction, text_direction, dim=1)).mean()
reg_tv = self.lambda_tv * self.get_image_prior_losses(target)
#total_loss = self.lambda_patch * loss_patch + self.content_weight * content_loss + reg_tv + self.lambda_dir * loss_glob
total_loss = self.content_weight * content_loss + reg_tv + self.lambda_dir * loss_glob
self.total_loss_epoch.append(total_loss)
self.optimizer.zero_grad()
total_loss.backward()
self.optimizer.step()
out_path = f'./process/{epoch}_{"exp0219"}_final.jpg'
target = torch.clamp(target, 0, 1)
save_image(target, out_path, nrow=1, normalize=True)
for name, param in self.style_net.named_parameters():
if param.grad is not None:
print(f"Gradients for {name}: {param.grad.mean()}")
else:
print(f"No gradient for {name}")
if epoch % 20 == 0:
print("After %d criterions:" % epoch)
print('Total loss: ', total_loss.item())
print('Content loss: ', content_loss.item())
print('patch loss: ', loss_patch.item())
print('dir loss: ', loss_glob.item())
print('TV loss: ', reg_tv.item())
return target.detach()
I want to replace the image editing part of InstructN2N, which uses InstructPix2Pix, with another model, CLIPStyler, that requires 200 iterations for training. My CLIPStyler works fine when run independently, but after replacing InstructPix2Pix, the output of CLIPStyler is the same for every iteration, even though the gradients are calculated successfully and the model parameters are changing. I don’t understand why the output is the same every time. This is preventing me from successfully editing the image in InstructN2N. Here is my CLIPStyler code used to replace InstructPix2Pix in the source files.
`from PIL import Image
import numpy as np
import os
import torch
import torch.nn
import torch.optim as optim
from torchvision import transforms, models
from torchvision.utils import save_image
import in2n.StyleNet
import in2n.utils as utils
import clip
import torch.nn.functional as F
from in2n.template import imagenet_templates
from typing import Union
from PIL import Image
import PIL
from torchvision import utils as vutils
from torchvision.transforms.functional import adjust_contrast
from rich.console import Console
from torchvision.models import VGG19_Weights
CONSOLE = Console(width=120)
class CLIPstyler:
def init(self, device: Union[torch.device, str],text):
`