File size: 7,963 Bytes
f0395ef
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
# Copyright (c) 2025 Hansheng Chen

import torch
import inspect

from copy import deepcopy
from mmgen.models.builder import MODELS, build_module

from .base_diffusion import BaseDiffusion
from lakonlab.utils import rgetattr


@MODELS.register_module()
class LatentDiffusionTextImage(BaseDiffusion):

    def __init__(self,
                 *args,
                 vae=None,
                 text_encoder=None,
                 **kwargs):
        super().__init__(*args, **kwargs)
        self.vae = build_module(vae) if vae is not None else None
        self.text_encoder = build_module(text_encoder) if text_encoder is not None else None

    def _prepare_train_minibatch_diffusion_args(self, data):
        if 'prompt_embed_kwargs' in data:
            prompt_embed_kwargs = data['prompt_embed_kwargs']
        elif 'prompt_kwargs' in data:
            assert self.text_encoder is not None, 'Text encoder must be provided for encoding text to embeddings.'
            prompt_embed_kwargs = self.text_encoder(**data['prompt_kwargs'])
        else:
            raise ValueError('Either `prompt_embed_kwargs` or `prompt_kwargs` should be provided in the input data.')

        if 'latents' in data:
            latents = data['latents']
        elif 'images' in data:
            assert self.vae is not None, 'VAE must be provided for encoding images to latents.'
            with torch.no_grad():
                if hasattr(self.vae, 'dtype'):
                    vae_dtype = self.vae.dtype
                else:
                    vae_dtype = next(self.vae.parameters()).dtype
                latents = self.vae.encode((data['images'] * 2 - 1).to(vae_dtype)).float()
        else:
            raise ValueError('Either `latents` or `images` should be provided in the input data.')

        v = next(iter(prompt_embed_kwargs.values()))
        bs = v.size(0)
        device = v.device

        diffusion_args = (self.patchify(latents), )
        diffusion_kwargs = prompt_embed_kwargs.copy()

        distilled_guidance_scale = self.train_cfg.get('distilled_guidance_scale', None)
        if distilled_guidance_scale is not None:
            distilled_guidance_scale = torch.full(
                (bs,), distilled_guidance_scale, dtype=torch.float32, device=device)
            diffusion_kwargs.update(guidance=distilled_guidance_scale)

        return diffusion_args, diffusion_kwargs, prompt_embed_kwargs, bs, device

    def _prepare_train_minibatch_teacher_args(self, data, prompt_embed_kwargs, bs, device):
        teacher_guidance_scale = self.train_cfg.get('teacher_guidance_scale', None)
        teacher_use_guidance = (teacher_guidance_scale is not None
                                and teacher_guidance_scale != 0.0 and teacher_guidance_scale != 1.0)
        if teacher_use_guidance:
            if 'negative_prompt_embed_kwargs' in data:
                negative_prompt_embed_kwargs = data['negative_prompt_embed_kwargs']
            elif 'negative_prompt_kwargs' in data:
                negative_prompt_embed_kwargs = self.text_encoder(**data['negative_prompt_kwargs'])
            else:
                raise ValueError(
                    'Either `negative_prompt_embed_kwargs` or `negative_prompt_kwargs` should be provided in the '
                    'input data for classifier-free guidance.')
            teacher_kwargs = {
                k: torch.cat([negative_prompt_embed_kwargs[k], v], dim=0)
                for k, v in prompt_embed_kwargs.items()}
            teacher_kwargs.update(guidance_scale=teacher_guidance_scale)
        else:
            teacher_kwargs = prompt_embed_kwargs.copy()

        teacher_distilled_guidance_scale = self.train_cfg.get('teacher_distilled_guidance_scale', None)
        if teacher_distilled_guidance_scale is not None:
            teacher_distilled_guidance_scale = torch.full(
                (bs * 2,) if teacher_use_guidance else (bs,),
                teacher_distilled_guidance_scale, dtype=torch.float32, device=device)
            teacher_kwargs.update(guidance=teacher_distilled_guidance_scale)

        return teacher_kwargs

    def _prepare_train_minibatch_args(self, data, running_status=None):
        diffusion_args, diffusion_kwargs, prompt_embed_kwargs, bs, device = \
            self._prepare_train_minibatch_diffusion_args(data)
        parameters = inspect.signature(rgetattr(self.diffusion, 'forward_train')).parameters
        if 'running_status' in parameters:
            diffusion_kwargs['running_status'] = running_status

        if 'teacher' in parameters and 'teacher_kwargs' in parameters and self.teacher is not None:
            teacher_kwargs = self._prepare_train_minibatch_teacher_args(
                data, prompt_embed_kwargs, bs, device)

            diffusion_kwargs.update(
                teacher=self.teacher,
                teacher_kwargs=teacher_kwargs)

        return bs, diffusion_args, diffusion_kwargs

    def val_step(self, data, test_cfg_override=dict(), **kwargs):
        if 'prompt_embed_kwargs' in data:
            prompt_embed_kwargs = data['prompt_embed_kwargs']
        elif 'prompt_kwargs' in data:
            assert self.text_encoder is not None, 'Text encoder must be provided for encoding text to embeddings.'
            prompt_embed_kwargs = self.text_encoder(**data['prompt_kwargs'])
        else:
            raise ValueError('Either `prompt_embed_kwargs` or `prompt_kwargs` should be provided in the input data.')

        v = next(iter(prompt_embed_kwargs.values()))
        bs = v.size(0)
        device = v.device

        cfg = deepcopy(self.test_cfg)
        cfg.update(test_cfg_override)
        guidance_scale = cfg.get('guidance_scale', 1.0)
        diffusion = self.diffusion_ema if self.diffusion_use_ema else self.diffusion

        with torch.no_grad():
            use_guidance = guidance_scale != 0.0 and guidance_scale != 1.0
            if use_guidance:
                if 'negative_prompt_embed_kwargs' in data:
                    negative_prompt_embed_kwargs = data['negative_prompt_embed_kwargs']
                elif 'negative_prompt_kwargs' in data:
                    negative_prompt_embed_kwargs = self.text_encoder(**data['negative_prompt_kwargs'])
                else:
                    raise ValueError(
                        'Either `negative_prompt_embed_kwargs` or `negative_prompt_kwargs` should be provided in the '
                        'input data for classifier-free guidance.')
                kwargs = {
                    k: torch.cat([negative_prompt_embed_kwargs[k], v], dim=0)
                    for k, v in prompt_embed_kwargs.items()}
            else:
                kwargs = prompt_embed_kwargs.copy()
            distilled_guidance_scale = cfg.get('distilled_guidance_scale', None)
            if distilled_guidance_scale is not None:
                distilled_guidance_scale = torch.full(
                    (bs * 2,) if use_guidance else (bs,),
                    distilled_guidance_scale, dtype=torch.float32, device=device)
                kwargs.update(guidance=distilled_guidance_scale)

            if 'noise' in data:
                noise = data['noise']
            else:
                latent_size = cfg['latent_size']
                noise = torch.randn((bs, *latent_size), device=device)
            noise = self.patchify(noise)
            latents_out = diffusion(
                noise=noise,
                guidance_scale=guidance_scale,
                test_cfg_override=test_cfg_override,
                **kwargs)
            latents_out = self.unpatchify(latents_out)

            if hasattr(self.vae, 'dtype'):
                vae_dtype = self.vae.dtype
            else:
                vae_dtype = next(self.vae.parameters()).dtype
            latents_out = latents_out.to(vae_dtype)

            out_images = (self.vae.decode(latents_out).float() / 2 + 0.5).clamp(min=0, max=1)

            return dict(num_samples=bs, pred_imgs=out_images)