File size: 13,973 Bytes
c8ff942
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
import torch
import torch.nn.functional as F


# Based on code from https://github.com/pix2pixzero/pix2pix-zero
def noise_regularization(
    e_t,
    noise_pred_optimal,
    lambda_kl,
    lambda_ac,
    num_reg_steps,
    num_ac_rolls,
    generator=None,
):
    for _outer in range(num_reg_steps):
        if lambda_kl > 0:
            _var = torch.autograd.Variable(e_t.detach().clone(), requires_grad=True)
            l_kld = patchify_latents_kl_divergence(_var, noise_pred_optimal)
            l_kld.backward()
            _grad = _var.grad.detach()
            _grad = torch.clip(_grad, -100, 100)
            e_t = e_t - lambda_kl * _grad
        if lambda_ac > 0:
            for _inner in range(num_ac_rolls):
                _var = torch.autograd.Variable(e_t.detach().clone(), requires_grad=True)
                l_ac = auto_corr_loss(_var, generator=generator)
                l_ac.backward()
                _grad = _var.grad.detach() / num_ac_rolls
                e_t = e_t - lambda_ac * _grad
        e_t = e_t.detach()

    return e_t


# Based on code from https://github.com/pix2pixzero/pix2pix-zero
def auto_corr_loss(x, random_shift=True, generator=None):
    B, C, H, W = x.shape
    assert B == 1
    x = x.squeeze(0)
    # x must be shape [C,H,W] now
    reg_loss = 0.0
    for ch_idx in range(x.shape[0]):
        noise = x[ch_idx][None, None, :, :]
        while True:
            if random_shift:
                roll_amount = torch.randint(
                    0, noise.shape[2] // 2, (1,), generator=generator
                ).item()
            else:
                roll_amount = 1
            reg_loss += (
                noise * torch.roll(noise, shifts=roll_amount, dims=2)
            ).mean() ** 2
            reg_loss += (
                noise * torch.roll(noise, shifts=roll_amount, dims=3)
            ).mean() ** 2
            if noise.shape[2] <= 8:
                break
            noise = F.avg_pool2d(noise, kernel_size=2)
    return reg_loss


def patchify_latents_kl_divergence(x0, x1, patch_size=4, num_channels=4):

    def patchify_tensor(input_tensor):
        patches = (
            input_tensor.unfold(1, patch_size, patch_size)
            .unfold(2, patch_size, patch_size)
            .unfold(3, patch_size, patch_size)
        )
        patches = patches.contiguous().view(-1, num_channels, patch_size, patch_size)
        return patches

    x0 = patchify_tensor(x0)
    x1 = patchify_tensor(x1)

    kl = latents_kl_divergence(x0, x1).sum()
    return kl


def latents_kl_divergence(x0, x1):
    EPSILON = 1e-6
    x0 = x0.view(x0.shape[0], x0.shape[1], -1)
    x1 = x1.view(x1.shape[0], x1.shape[1], -1)
    mu0 = x0.mean(dim=-1)
    mu1 = x1.mean(dim=-1)
    var0 = x0.var(dim=-1)
    var1 = x1.var(dim=-1)
    kl = (
        torch.log((var1 + EPSILON) / (var0 + EPSILON))
        + (var0 + (mu0 - mu1) ** 2) / (var1 + EPSILON)
        - 1
    )
    kl = torch.abs(kl).sum(dim=-1)
    return kl


def inversion_step(
    pipe,
    z_t: torch.tensor,
    t: torch.tensor,
    prompt_embeds,
    added_cond_kwargs,
    num_gd_steps: int = 3,
    gd_step_size: float = 0.001,
    optimization_start: int = 0,
    normalize: bool = False,
    use_cfgpp: bool = False,
    generator=None,
) -> torch.tensor:
    extra_step_kwargs = {}
    approximated_z_tp1 = z_t
    recon_diff = None
    # with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
    noise_pred = unet_pass(
        pipe, approximated_z_tp1, t, prompt_embeds, added_cond_kwargs
    )
    # perform guidance
    if pipe.do_classifier_free_guidance:
        noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
        noise_pred = noise_pred_uncond + pipe.guidance_scale * (
            noise_pred_text - noise_pred_uncond
        )

    if use_cfgpp:
        extra_step_kwargs["noise_pred_uncond"] = noise_pred_uncond
    approximated_z_tp1 = pipe.scheduler.inv_step(
        noise_pred, t, z_t, **extra_step_kwargs, return_dict=False
    )[0].detach()

    numel_sqrt = approximated_z_tp1.numel() ** 0.5
    total_steps = 1000
    if t < optimization_start:
        return approximated_z_tp1, recon_diff
    if normalize:
        alpha = (t - optimization_start) / (total_steps - optimization_start)
        alpha = alpha**0.1
        approximated_z_tp1 = approximated_z_tp1 / (
            (alpha * torch.linalg.vector_norm(approximated_z_tp1) / numel_sqrt)
            + (1 - alpha)
        )
        if t < 100:
            num_gd_steps += 5
    approximated_z_tp1.requires_grad = True
    optimizer = torch.optim.Adam([approximated_z_tp1], lr=gd_step_size, eps=1e-8)
    for i in range(num_gd_steps + 1):
        # uncomment in order to skip reporting the post optimization reconstruction loss
        if i == num_gd_steps:
            break
        noise_pred = unet_pass(
            pipe, approximated_z_tp1, t, prompt_embeds, added_cond_kwargs
        )
        # perform guidance
        if pipe.do_classifier_free_guidance:
            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
            noise_pred = noise_pred_uncond + pipe.guidance_scale * (
                noise_pred_text - noise_pred_uncond
            )
        if use_cfgpp:
            extra_step_kwargs["noise_pred_uncond"] = noise_pred_uncond
        approximated_z_t = pipe.scheduler.step(
            noise_pred, t, approximated_z_tp1, **extra_step_kwargs, return_dict=False
        )[0]
        optimizer.zero_grad()
        norm = torch.linalg.vector_norm(approximated_z_t) / numel_sqrt
        recon_mse_loss = torch.nn.functional.mse_loss(approximated_z_t, z_t)
        loss = recon_mse_loss
        if i == num_gd_steps:
            recon_diff = approximated_z_t - z_t
            print(
                f"t={t} Post optimization Reconstruction Loss: {recon_mse_loss} Norm: {norm}"
            )
            break
        print(f"t={t} round={i} Reconstruction Loss: {recon_mse_loss} Norm: {norm}")
        loss.backward()
        # approximated_z_tp1.grad = torch.where(
        #     approximated_z_tp1 > 0,
        #     torch.clamp(approximated_z_tp1.grad, min=0),
        #     torch.clamp(approximated_z_tp1.grad, max=0),
        # )
        # print(f"approximated_z_tp1: {approximated_z_tp1}")
        # print(f"Gradient: {approximated_z_tp1.grad}")
        # grad_norm = approximated_z_tp1.grad.norm().item()
        # print(f"Gradient norm: {grad_norm}")
        optimizer.step()
        # print(f"approximated_z_tp1 after step: {approximated_z_tp1}")

    return approximated_z_tp1, recon_diff

def renoise_inversion_step(
    pipe,
    z_t: torch.tensor,
    t: torch.tensor,
    prompt_embeds,
    added_cond_kwargs,
    generator=None,
    use_cfgpp: bool = False,
) -> torch.tensor:
    extra_step_kwargs = {}
    avg_range = pipe.cfg.renoise_config.average_first_step_range if t.item() < pipe.cfg.renoise_config.renoise_first_step_max_timestep else pipe.cfg.renoise_config.average_step_range
    num_renoise_steps = min(pipe.cfg.renoise_config.max_num_renoise_steps_first_step, pipe.cfg.renoise_config.num_renoise_steps) if t.item() < pipe.cfg.renoise_config.renoise_first_step_max_timestep else pipe.cfg.renoise_config.num_renoise_steps

    noise_pred_avg = None
    noise_pred_optimal = None
    noise_pred_avg_uncond = None
    noise_pred_optimal_uncond = None
    z_tp1_forward = pipe.scheduler.add_noise(pipe.z_0, pipe.noise, t.view((1))).detach()

    approximated_z_tp1 = z_t.clone()
    for i in range(num_renoise_steps + 1):

        with torch.no_grad():
            # if noise regularization is enabled, we need to double the batch size for the first step
            if pipe.cfg.renoise_config.noise_regularization_num_reg_steps > 0 and i == 0:
                approximated_z_tp1 = torch.cat([z_tp1_forward, approximated_z_tp1])
                prompt_embeds_in = torch.cat([prompt_embeds, prompt_embeds])
                if added_cond_kwargs is not None:
                    added_cond_kwargs_in = {}
                    added_cond_kwargs_in['text_embeds'] = torch.cat([added_cond_kwargs['text_embeds'], added_cond_kwargs['text_embeds']])
                    added_cond_kwargs_in['time_ids'] = torch.cat([added_cond_kwargs['time_ids'], added_cond_kwargs['time_ids']])
                    if "image_embeds" in added_cond_kwargs.keys(): 
                        print("image_embeds:")
                        print(added_cond_kwargs["image_embeds"])
                        added_cond_kwargs_in["image_embeds"] = [torch.cat([added_cond_kwargs["image_embeds"][0], added_cond_kwargs["image_embeds"][0]])]
                        added_cond_kwargs_in["image_embeds"] = added_cond_kwargs["image_embeds"] # maybe this one is correct
                else:
                    added_cond_kwargs_in = None
            else:
                prompt_embeds_in = prompt_embeds
                added_cond_kwargs_in = added_cond_kwargs

            noise_pred = unet_pass(
                pipe, 
                approximated_z_tp1, 
                t, 
                prompt_embeds_in, 
                added_cond_kwargs=added_cond_kwargs_in
            )

            # if noise regularization is enabled, we need to split the batch size for the first step
            if pipe.cfg.renoise_config.noise_regularization_num_reg_steps > 0 and i == 0:
                noise_pred_optimal, noise_pred = noise_pred.chunk(2)
                if pipe.do_classifier_free_guidance:
                    noise_pred_optimal_uncond, noise_pred_optimal_text = noise_pred_optimal.chunk(2)
                    noise_pred_optimal = noise_pred_optimal_uncond + pipe.guidance_scale * (noise_pred_optimal_text - noise_pred_optimal_uncond)
                noise_pred_optimal = noise_pred_optimal.detach()

            # perform guidance
            if pipe.do_classifier_free_guidance:
                noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
                noise_pred = noise_pred_uncond + pipe.guidance_scale * (noise_pred_text - noise_pred_uncond)

            # Calculate average noise
            if pipe.cfg.renoise_config.average_latent_estimations and i >= avg_range[0] and i < avg_range[1]:
                j = i - avg_range[0]
                if noise_pred_avg is None:
                    noise_pred_avg = noise_pred.clone()
                else:
                    noise_pred_avg = j * noise_pred_avg / (j + 1) + noise_pred / (j + 1)
                if use_cfgpp:
                    if noise_pred_avg_uncond is None:
                        noise_pred_avg_uncond = noise_pred_uncond.clone()
                    else:
                        noise_pred_avg_uncond = j * noise_pred_avg_uncond / (j + 1) + noise_pred_uncond / (j + 1)

        if i >= avg_range[0] or (not pipe.cfg.renoise_config.average_latent_estimations and i > 0):
            noise_pred = noise_regularization(noise_pred, noise_pred_optimal, lambda_kl=pipe.cfg.renoise_config.noise_regularization_lambda_kl, lambda_ac=pipe.cfg.renoise_config.noise_regularization_lambda_ac, num_reg_steps=pipe.cfg.renoise_config.noise_regularization_num_reg_steps, num_ac_rolls=pipe.cfg.renoise_config.noise_regularization_num_ac_rolls, generator=generator)
        
        if use_cfgpp:
            extra_step_kwargs["noise_pred_uncond"] = noise_pred_uncond
        approximated_z_tp1 = pipe.scheduler.inv_step(noise_pred, t, z_t, **extra_step_kwargs, return_dict=False)[0].detach()

    # if average latents is enabled, we need to perform an additional step with the average noise
    if pipe.cfg.renoise_config.average_latent_estimations and noise_pred_avg is not None:
        noise_pred_avg = noise_regularization(noise_pred_avg, noise_pred_optimal, lambda_kl=pipe.cfg.renoise_config.noise_regularization_lambda_kl, lambda_ac=pipe.cfg.renoise_config.noise_regularization_lambda_ac, num_reg_steps=pipe.cfg.renoise_config.noise_regularization_num_reg_steps, num_ac_rolls=pipe.cfg.renoise_config.noise_regularization_num_ac_rolls, generator=generator)
        if use_cfgpp:
            noise_pred_avg_uncond = noise_regularization(noise_pred_avg_uncond, noise_pred_optimal_uncond, lambda_kl=pipe.cfg.renoise_config.noise_regularization_lambda_kl, lambda_ac=pipe.cfg.renoise_config.noise_regularization_lambda_ac, num_reg_steps=pipe.cfg.renoise_config.noise_regularization_num_reg_steps, num_ac_rolls=pipe.cfg.renoise_config.noise_regularization_num_ac_rolls, generator=generator)
            extra_step_kwargs["noise_pred_uncond"] = noise_pred_avg_uncond
        approximated_z_tp1 = pipe.scheduler.inv_step(noise_pred_avg, t, z_t, **extra_step_kwargs, return_dict=False)[0].detach()

    # perform noise correction
    if pipe.cfg.renoise_config.perform_noise_correction:
        noise_pred = unet_pass(
            pipe, 
            approximated_z_tp1, 
            t, 
            prompt_embeds, 
            added_cond_kwargs=added_cond_kwargs,
        )

        # perform guidance
        if pipe.do_classifier_free_guidance:
            noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
            noise_pred = noise_pred_uncond + pipe.guidance_scale * (noise_pred_text - noise_pred_uncond)
        
        if use_cfgpp:
            extra_step_kwargs["noise_pred_uncond"] = noise_pred_uncond
        pipe.scheduler.step_and_update_noise(noise_pred, t, approximated_z_tp1, z_t, return_dict=False, optimize_epsilon_type=pipe.cfg.renoise_config.perform_noise_correction)

    return approximated_z_tp1


@torch.no_grad()
def unet_pass(pipe, z_t, t, prompt_embeds, added_cond_kwargs):
    latent_model_input = (
        torch.cat([z_t] * 2) if pipe.do_classifier_free_guidance else z_t
    )
    latent_model_input = pipe.scheduler.scale_model_input(latent_model_input, t)
    return pipe.unet(
        latent_model_input,
        t,
        encoder_hidden_states=prompt_embeds,
        timestep_cond=None,
        cross_attention_kwargs=pipe.cross_attention_kwargs,
        added_cond_kwargs=added_cond_kwargs,
        return_dict=False,
    )[0]