Skip to content

How to replace the ip2p with Clipstyler? #104

Description

@Jillianbug

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):

    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()

`

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions