File size: 6,181 Bytes
5a14c00
ff63d1b
 
 
 
5a14c00
ff63d1b
5a14c00
ff63d1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5a14c00
ff63d1b
 
 
 
 
 
 
 
 
5a14c00
ff63d1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5a14c00
 
ff63d1b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
159
160
161
162
163
164
165
166
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

clip_model = None
tokenizer = None
vgg19 = None
cnn_model = None

def load_models():
    """Load models once and cache them"""
    global clip_model, tokenizer, vgg19, cnn_model
    
    if clip_model is None:
        clip_model = CLIPModel.from_pretrained("openai/clip-vit-base-patch32").to(DEVICE)
    
    if tokenizer is None:
        tokenizer = AutoTokenizer.from_pretrained("openai/clip-vit-base-patch32")
    
    if vgg19 is None:
        vgg19 = models.vgg19(pretrained=True).features.to(DEVICE)
        for param in vgg19.parameters():
            param.requires_grad_(False)
    
    if cnn_model is None:
        cnn_model = UNet(text_dim=512).to(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, num_steps=NUM_EPOCHS, 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()
    
    # Reset CNN model for fresh start on each image
    reset_cnn_model()
    
    # image features
    features = get_features(normalize(img), vgg19)

    # cnn model
    #cnn_model = UNet().to(DEVICE)

    # load ADAM
    optimizer = optim.Adam(cnn_model.parameters(), lr=LR)
    scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=100, gamma=0.5)

    # data augmentation
    crop = transforms.Compose([transforms.RandomCrop(CROP_SIZE)])
    augment = transforms.Compose([transforms.RandomPerspective(fill=0, p=1, distortion_scale=0.5), transforms.Resize(RESIZE)])

    # initialize variables
    content_loss_epoch = []
    style_loss_epoch = []
    total_loss_epoch = []
    output_image = img
    mean_img = torch.mean(img, dim=(2, 3), keepdim=False).squeeze(0)
    mean_img = [mean_img[0].item(), mean_img[1].item(), mean_img[2].item()]
    target = img.clone().requires_grad_(True).to(DEVICE)

    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)

        edited_source = prompt_ensemble(source)
        source_tokens = tokenizer(edited_source, padding=True, return_tensors="pt").to(DEVICE)
        source_features = clip_model.get_text_features(**source_tokens)
        source_features = source_features.mean(axis=0, keepdim=True)
        source_features /= source_features.norm(dim=-1, keepdim=True)

        img_features = clip_model.get_image_features(pixel_values=clip_normalize(img))
        img_features /= (img_features.clone().norm(dim=-1, keepdim=True))

    for epoch in range(num_steps + 1):
        scheduler.step()
        #target = cnn_model(img).requires_grad_(True).to(DEVICE)
        target = cnn_model(img, text_embedding=text_features).to(DEVICE)

        target_features = get_features(normalize(target), vgg19)

        content_loss = 0.0
        content_loss += torch.mean((target_features['conv4_2'] - features['conv4_2']) ** 2)
        content_loss += torch.mean((target_features['conv5_2'] - features['conv5_2']) ** 2)

        loss_patch = 0 
        img_proc = []
        for n in range(NUM_CROPS):
            target_crop = crop(target)
            target_crop = augment(target_crop)
            img_proc.append(target_crop)

        img_proc = torch.cat(img_proc, dim=0)
        img_aug = img_proc

        img_aug_features = clip_model.get_image_features(pixel_values=clip_normalize(img_aug))
        img_aug_features /= (img_aug_features.clone().norm(dim=-1, keepdim=True))

        img_direction = img_aug_features - img_features
        img_direction /= img_direction.clone().norm(dim=-1, keepdim=True)

        text_direction = (text_features - source_features).repeat(img_aug_features.size(0), 1)
        text_direction /= text_direction.norm(dim=-1, keepdim=True)

        loss_calc = (1 - torch.cosine_similarity(img_direction, text_direction, dim=1))
        loss_calc[loss_calc < PATCH_THRESHOLD] = 0.0        
        loss_patch += loss_calc.mean()

        glob_features = clip_model.get_image_features(pixel_values=clip_normalize(target))
        glob_features /= (glob_features.clone().norm(dim=-1, keepdim=True))
        glob_direction = glob_features - img_features
        glob_direction /= glob_direction.clone().norm(dim=-1, keepdim=True)

        loss_glob = (1 - torch.cosine_similarity(glob_direction, text_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
        total_loss_epoch.append(total_loss)

        optimizer.zero_grad()
        total_loss.backward()
        optimizer.step()
    
    # 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