comdoleger commited on
Commit
ebde842
·
verified ·
1 Parent(s): 9e365dd

Upload extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/flux_kontext/flux_kontext.py ADDED
@@ -0,0 +1,420 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from typing import TYPE_CHECKING, List
3
+
4
+ import torch
5
+ import torchvision
6
+ import yaml
7
+ from toolkit import train_tools
8
+ from toolkit.config_modules import GenerateImageConfig, ModelConfig
9
+ from PIL import Image
10
+ from toolkit.models.base_model import BaseModel
11
+ from diffusers import FluxTransformer2DModel, AutoencoderKL, FluxKontextPipeline
12
+ from toolkit.basic import flush
13
+ from toolkit.prompt_utils import PromptEmbeds
14
+ from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
15
+ from toolkit.models.flux import add_model_gpu_splitter_to_flux, bypass_flux_guidance, restore_flux_guidance
16
+ from toolkit.dequantize import patch_dequantization_on_save
17
+ from toolkit.accelerator import get_accelerator, unwrap_model
18
+ from optimum.quanto import freeze, QTensor
19
+ from toolkit.util.mask import generate_random_mask, random_dialate_mask
20
+ from toolkit.util.quantize import quantize, get_qtype
21
+ from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer
22
+ from einops import rearrange, repeat
23
+ import random
24
+ import torch.nn.functional as F
25
+
26
+ if TYPE_CHECKING:
27
+ from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
28
+
29
+ scheduler_config = {
30
+ "base_image_seq_len": 256,
31
+ "base_shift": 0.5,
32
+ "max_image_seq_len": 4096,
33
+ "max_shift": 1.15,
34
+ "num_train_timesteps": 1000,
35
+ "shift": 3.0,
36
+ "use_dynamic_shifting": True
37
+ }
38
+
39
+
40
+
41
+ class FluxKontextModel(BaseModel):
42
+ arch = "flux_kontext"
43
+
44
+ def __init__(
45
+ self,
46
+ device,
47
+ model_config: ModelConfig,
48
+ dtype='bf16',
49
+ custom_pipeline=None,
50
+ noise_scheduler=None,
51
+ **kwargs
52
+ ):
53
+ super().__init__(
54
+ device,
55
+ model_config,
56
+ dtype,
57
+ custom_pipeline,
58
+ noise_scheduler,
59
+ **kwargs
60
+ )
61
+ self.is_flow_matching = True
62
+ self.is_transformer = True
63
+ self.target_lora_modules = ['FluxTransformer2DModel']
64
+
65
+ # static method to get the noise scheduler
66
+ @staticmethod
67
+ def get_train_scheduler():
68
+ return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
69
+
70
+ def get_bucket_divisibility(self):
71
+ return 16
72
+
73
+ def load_model(self):
74
+ dtype = self.torch_dtype
75
+ self.print_and_status_update("Loading Flux Kontext model")
76
+ # will be updated if we detect a existing checkpoint in training folder
77
+ model_path = self.model_config.name_or_path
78
+ # this is the original path put in the model directory
79
+ # it is here because for finetuning we only save the transformer usually
80
+ # so we need this for the VAE, te, etc
81
+ base_model_path = self.model_config.extras_name_or_path
82
+
83
+ transformer_path = model_path
84
+ transformer_subfolder = 'transformer'
85
+ if os.path.exists(transformer_path):
86
+ transformer_subfolder = None
87
+ transformer_path = os.path.join(transformer_path, 'transformer')
88
+ # check if the path is a full checkpoint.
89
+ te_folder_path = os.path.join(model_path, 'text_encoder')
90
+ # if we have the te, this folder is a full checkpoint, use it as the base
91
+ if os.path.exists(te_folder_path):
92
+ base_model_path = model_path
93
+
94
+ self.print_and_status_update("Loading transformer")
95
+ transformer = FluxTransformer2DModel.from_pretrained(
96
+ transformer_path,
97
+ subfolder=transformer_subfolder,
98
+ torch_dtype=dtype
99
+ )
100
+ transformer.to(self.quantize_device, dtype=dtype)
101
+
102
+ if self.model_config.quantize:
103
+ # patch the state dict method
104
+ patch_dequantization_on_save(transformer)
105
+ quantization_type = get_qtype(self.model_config.qtype)
106
+ self.print_and_status_update("Quantizing transformer")
107
+ quantize(transformer, weights=quantization_type,
108
+ **self.model_config.quantize_kwargs)
109
+ freeze(transformer)
110
+ transformer.to(self.device_torch)
111
+ else:
112
+ transformer.to(self.device_torch, dtype=dtype)
113
+
114
+ flush()
115
+
116
+ self.print_and_status_update("Loading T5")
117
+ tokenizer_2 = T5TokenizerFast.from_pretrained(
118
+ base_model_path, subfolder="tokenizer_2", torch_dtype=dtype
119
+ )
120
+ text_encoder_2 = T5EncoderModel.from_pretrained(
121
+ base_model_path, subfolder="text_encoder_2", torch_dtype=dtype
122
+ )
123
+ text_encoder_2.to(self.device_torch, dtype=dtype)
124
+ flush()
125
+
126
+ if self.model_config.quantize_te:
127
+ self.print_and_status_update("Quantizing T5")
128
+ quantize(text_encoder_2, weights=get_qtype(
129
+ self.model_config.qtype))
130
+ freeze(text_encoder_2)
131
+ flush()
132
+
133
+ self.print_and_status_update("Loading CLIP")
134
+ text_encoder = CLIPTextModel.from_pretrained(
135
+ base_model_path, subfolder="text_encoder", torch_dtype=dtype)
136
+ tokenizer = CLIPTokenizer.from_pretrained(
137
+ base_model_path, subfolder="tokenizer", torch_dtype=dtype)
138
+ text_encoder.to(self.device_torch, dtype=dtype)
139
+
140
+ self.print_and_status_update("Loading VAE")
141
+ vae = AutoencoderKL.from_pretrained(
142
+ base_model_path, subfolder="vae", torch_dtype=dtype)
143
+
144
+ self.noise_scheduler = FluxKontextModel.get_train_scheduler()
145
+
146
+ self.print_and_status_update("Making pipe")
147
+
148
+ pipe: FluxKontextPipeline = FluxKontextPipeline(
149
+ scheduler=self.noise_scheduler,
150
+ text_encoder=text_encoder,
151
+ tokenizer=tokenizer,
152
+ text_encoder_2=None,
153
+ tokenizer_2=tokenizer_2,
154
+ vae=vae,
155
+ transformer=None,
156
+ )
157
+ # for quantization, it works best to do these after making the pipe
158
+ pipe.text_encoder_2 = text_encoder_2
159
+ pipe.transformer = transformer
160
+
161
+ self.print_and_status_update("Preparing Model")
162
+
163
+ text_encoder = [pipe.text_encoder, pipe.text_encoder_2]
164
+ tokenizer = [pipe.tokenizer, pipe.tokenizer_2]
165
+
166
+ pipe.transformer = pipe.transformer.to(self.device_torch)
167
+
168
+ flush()
169
+ # just to make sure everything is on the right device and dtype
170
+ text_encoder[0].to(self.device_torch)
171
+ text_encoder[0].requires_grad_(False)
172
+ text_encoder[0].eval()
173
+ text_encoder[1].to(self.device_torch)
174
+ text_encoder[1].requires_grad_(False)
175
+ text_encoder[1].eval()
176
+ pipe.transformer = pipe.transformer.to(self.device_torch)
177
+ flush()
178
+
179
+ # save it to the model class
180
+ self.vae = vae
181
+ self.text_encoder = text_encoder # list of text encoders
182
+ self.tokenizer = tokenizer # list of tokenizers
183
+ self.model = pipe.transformer
184
+ self.pipeline = pipe
185
+ self.print_and_status_update("Model Loaded")
186
+
187
+ def get_generation_pipeline(self):
188
+ scheduler = FluxKontextModel.get_train_scheduler()
189
+
190
+ pipeline: FluxKontextPipeline = FluxKontextPipeline(
191
+ scheduler=scheduler,
192
+ text_encoder=unwrap_model(self.text_encoder[0]),
193
+ tokenizer=self.tokenizer[0],
194
+ text_encoder_2=unwrap_model(self.text_encoder[1]),
195
+ tokenizer_2=self.tokenizer[1],
196
+ vae=unwrap_model(self.vae),
197
+ transformer=unwrap_model(self.transformer)
198
+ )
199
+
200
+ pipeline = pipeline.to(self.device_torch)
201
+
202
+ return pipeline
203
+
204
+ def generate_single_image(
205
+ self,
206
+ pipeline: FluxKontextPipeline,
207
+ gen_config: GenerateImageConfig,
208
+ conditional_embeds: PromptEmbeds,
209
+ unconditional_embeds: PromptEmbeds,
210
+ generator: torch.Generator,
211
+ extra: dict,
212
+ ):
213
+ if gen_config.ctrl_img is None:
214
+ raise ValueError(
215
+ "Control image is required for Flux Kontext model generation."
216
+ )
217
+ else:
218
+ control_img = Image.open(gen_config.ctrl_img)
219
+ control_img = control_img.convert("RGB")
220
+ # resize to width and height
221
+ if control_img.size != (gen_config.width, gen_config.height):
222
+ control_img = control_img.resize(
223
+ (gen_config.width, gen_config.height), Image.BILINEAR
224
+ )
225
+ gen_config.width = int(gen_config.width // 16 * 16)
226
+ gen_config.height = int(gen_config.height // 16 * 16)
227
+ img = pipeline(
228
+ image=control_img,
229
+ prompt_embeds=conditional_embeds.text_embeds,
230
+ pooled_prompt_embeds=conditional_embeds.pooled_embeds,
231
+ height=gen_config.height,
232
+ width=gen_config.width,
233
+ num_inference_steps=gen_config.num_inference_steps,
234
+ guidance_scale=gen_config.guidance_scale,
235
+ latents=gen_config.latents,
236
+ generator=generator,
237
+ max_area=gen_config.height * gen_config.width,
238
+ _auto_resize=False,
239
+ **extra
240
+ ).images[0]
241
+ return img
242
+
243
+ def get_noise_prediction(
244
+ self,
245
+ latent_model_input: torch.Tensor,
246
+ timestep: torch.Tensor, # 0 to 1000 scale
247
+ text_embeddings: PromptEmbeds,
248
+ guidance_embedding_scale: float,
249
+ bypass_guidance_embedding: bool,
250
+ **kwargs
251
+ ):
252
+ with torch.no_grad():
253
+ bs, c, h, w = latent_model_input.shape
254
+ # if we have a control on the channel dimension, put it on the batch for packing
255
+ has_control = False
256
+ if latent_model_input.shape[1] == 32:
257
+ # chunk it and stack it on batch dimension
258
+ # dont update batch size for img_its
259
+ lat, control = torch.chunk(latent_model_input, 2, dim=1)
260
+ latent_model_input = torch.cat([lat, control], dim=0)
261
+ has_control = True
262
+
263
+ latent_model_input_packed = rearrange(
264
+ latent_model_input,
265
+ "b c (h ph) (w pw) -> b (h w) (c ph pw)",
266
+ ph=2,
267
+ pw=2
268
+ )
269
+
270
+ img_ids = torch.zeros(h // 2, w // 2, 3)
271
+ img_ids[..., 1] = img_ids[..., 1] + torch.arange(h // 2)[:, None]
272
+ img_ids[..., 2] = img_ids[..., 2] + torch.arange(w // 2)[None, :]
273
+ img_ids = repeat(img_ids, "h w c -> b (h w) c",
274
+ b=bs).to(self.device_torch)
275
+
276
+ # handle control image ids
277
+ if has_control:
278
+ ctrl_ids = img_ids.clone()
279
+ ctrl_ids[..., 0] = 1
280
+ img_ids = torch.cat([img_ids, ctrl_ids], dim=1)
281
+
282
+
283
+ txt_ids = torch.zeros(
284
+ bs, text_embeddings.text_embeds.shape[1], 3).to(self.device_torch)
285
+
286
+ # # handle guidance
287
+ if self.unet_unwrapped.config.guidance_embeds:
288
+ if isinstance(guidance_embedding_scale, list):
289
+ guidance = torch.tensor(
290
+ guidance_embedding_scale, device=self.device_torch)
291
+ else:
292
+ guidance = torch.tensor(
293
+ [guidance_embedding_scale], device=self.device_torch)
294
+ # Expand guidance to match original batch_size
295
+ guidance = guidance.expand(bs)
296
+ else:
297
+ guidance = None
298
+
299
+ if bypass_guidance_embedding:
300
+ bypass_flux_guidance(self.unet)
301
+
302
+ cast_dtype = self.unet.dtype
303
+ # changes from orig implementation
304
+ if txt_ids.ndim == 3:
305
+ txt_ids = txt_ids[0]
306
+ if img_ids.ndim == 3:
307
+ img_ids = img_ids[0]
308
+
309
+ latent_size = latent_model_input_packed.shape[1]
310
+ # move the kontext channels. We have them on batch dimension to here, but need to put them on the latent dimension
311
+ if has_control:
312
+ latent, control = torch.chunk(latent_model_input_packed, 2, dim=0)
313
+ latent_model_input_packed = torch.cat(
314
+ [latent, control], dim=1
315
+ )
316
+ latent_size = latent.shape[1]
317
+
318
+ noise_pred = self.unet(
319
+ hidden_states=latent_model_input_packed.to(
320
+ self.device_torch, cast_dtype),
321
+ timestep=timestep / 1000,
322
+ encoder_hidden_states=text_embeddings.text_embeds.to(
323
+ self.device_torch, cast_dtype),
324
+ pooled_projections=text_embeddings.pooled_embeds.to(
325
+ self.device_torch, cast_dtype),
326
+ txt_ids=txt_ids,
327
+ img_ids=img_ids,
328
+ guidance=guidance,
329
+ return_dict=False,
330
+ **kwargs,
331
+ )[0]
332
+
333
+ # remove kontext image conditioning
334
+ noise_pred = noise_pred[:, :latent_size]
335
+
336
+ if isinstance(noise_pred, QTensor):
337
+ noise_pred = noise_pred.dequantize()
338
+
339
+ noise_pred = rearrange(
340
+ noise_pred,
341
+ "b (h w) (c ph pw) -> b c (h ph) (w pw)",
342
+ h=latent_model_input.shape[2] // 2,
343
+ w=latent_model_input.shape[3] // 2,
344
+ ph=2,
345
+ pw=2,
346
+ c=self.vae.config.latent_channels
347
+ )
348
+
349
+ if bypass_guidance_embedding:
350
+ restore_flux_guidance(self.unet)
351
+
352
+ return noise_pred
353
+
354
+ def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
355
+ if self.pipeline.text_encoder.device != self.device_torch:
356
+ self.pipeline.text_encoder.to(self.device_torch)
357
+ prompt_embeds, pooled_prompt_embeds = train_tools.encode_prompts_flux(
358
+ self.tokenizer,
359
+ self.text_encoder,
360
+ prompt,
361
+ max_length=512,
362
+ )
363
+ pe = PromptEmbeds(
364
+ prompt_embeds
365
+ )
366
+ pe.pooled_embeds = pooled_prompt_embeds
367
+ return pe
368
+
369
+ def get_model_has_grad(self):
370
+ # return from a weight if it has grad
371
+ return self.model.proj_out.weight.requires_grad
372
+
373
+ def get_te_has_grad(self):
374
+ # return from a weight if it has grad
375
+ return self.text_encoder[1].encoder.block[0].layer[0].SelfAttention.q.weight.requires_grad
376
+
377
+ def save_model(self, output_path, meta, save_dtype):
378
+ # only save the unet
379
+ transformer: FluxTransformer2DModel = unwrap_model(self.model)
380
+ transformer.save_pretrained(
381
+ save_directory=os.path.join(output_path, 'transformer'),
382
+ safe_serialization=True,
383
+ )
384
+
385
+ meta_path = os.path.join(output_path, 'aitk_meta.yaml')
386
+ with open(meta_path, 'w') as f:
387
+ yaml.dump(meta, f)
388
+
389
+ def get_loss_target(self, *args, **kwargs):
390
+ noise = kwargs.get('noise')
391
+ batch = kwargs.get('batch')
392
+ return (noise - batch.latents).detach()
393
+
394
+ def condition_noisy_latents(self, latents: torch.Tensor, batch:'DataLoaderBatchDTO'):
395
+ with torch.no_grad():
396
+ control_tensor = batch.control_tensor
397
+ if control_tensor is not None:
398
+ self.vae.to(self.device_torch)
399
+ # we are not packed here, so we just need to pass them so we can pack them later
400
+ control_tensor = control_tensor * 2 - 1
401
+ control_tensor = control_tensor.to(self.vae_device_torch, dtype=self.torch_dtype)
402
+
403
+ # if it is not the size of batch.tensor, (bs,ch,h,w) then we need to resize it
404
+ if batch.tensor is not None:
405
+ target_h, target_w = batch.tensor.shape[2], batch.tensor.shape[3]
406
+ else:
407
+ # When caching latents, batch.tensor is None. We get the size from the file_items instead.
408
+ target_h = batch.file_items[0].crop_height
409
+ target_w = batch.file_items[0].crop_width
410
+
411
+ if control_tensor.shape[2] != target_h or control_tensor.shape[3] != target_w:
412
+ control_tensor = F.interpolate(control_tensor, size=(target_h, target_w), mode='bilinear')
413
+
414
+ control_latent = self.encode_images(control_tensor).to(latents.device, latents.dtype)
415
+ latents = torch.cat((latents, control_latent), dim=1)
416
+
417
+ return latents.detach()
418
+
419
+ def get_base_model_version(self):
420
+ return "flux.1_kontext"