comdoleger commited on
Commit
9e1d1e6
·
verified ·
1 Parent(s): 28cfe7a

Upload extensions_built_in/diffusion_models/hidream/hidream_model.py with huggingface_hub

Browse files
extensions_built_in/diffusion_models/hidream/hidream_model.py ADDED
@@ -0,0 +1,453 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ from typing import TYPE_CHECKING, List, Optional
3
+
4
+ import einops
5
+ import torch
6
+ import torchvision
7
+ import yaml
8
+ from toolkit import train_tools
9
+ from toolkit.config_modules import GenerateImageConfig, ModelConfig
10
+ from PIL import Image
11
+ from toolkit.models.base_model import BaseModel
12
+ from diffusers import AutoencoderKL, TorchAoConfig
13
+ from toolkit.basic import flush
14
+ from toolkit.prompt_utils import PromptEmbeds
15
+ from toolkit.samplers.custom_flowmatch_sampler import CustomFlowMatchEulerDiscreteScheduler
16
+ from toolkit.models.flux import add_model_gpu_splitter_to_flux, bypass_flux_guidance, restore_flux_guidance
17
+ from toolkit.dequantize import patch_dequantization_on_save
18
+ from toolkit.accelerator import get_accelerator, unwrap_model
19
+ from optimum.quanto import freeze, QTensor
20
+ from toolkit.util.mask import generate_random_mask, random_dialate_mask
21
+ from toolkit.util.quantize import quantize, get_qtype
22
+ from transformers import T5TokenizerFast, T5EncoderModel, CLIPTextModel, CLIPTokenizer, TorchAoConfig as TorchAoConfigTransformers
23
+ from .src.pipelines.hidream_image.pipeline_hidream_image import HiDreamImagePipeline
24
+ from .src.models.transformers.transformer_hidream_image import HiDreamImageTransformer2DModel
25
+ from .src.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler
26
+ from transformers import LlamaForCausalLM, PreTrainedTokenizerFast
27
+ from einops import rearrange, repeat
28
+ import random
29
+ import torch.nn.functional as F
30
+ from tqdm import tqdm
31
+ from transformers import (
32
+ CLIPTextModelWithProjection,
33
+ CLIPTokenizer,
34
+ T5EncoderModel,
35
+ T5Tokenizer,
36
+ LlamaForCausalLM,
37
+ PreTrainedTokenizerFast
38
+ )
39
+
40
+ if TYPE_CHECKING:
41
+ from toolkit.data_transfer_object.data_loader import DataLoaderBatchDTO
42
+
43
+ scheduler_config = {
44
+ "num_train_timesteps": 1000,
45
+ "shift": 3.0
46
+ }
47
+
48
+ # LLAMA_MODEL_NAME = "meta-llama/Meta-Llama-3.1-8B-Instruct"
49
+ LLAMA_MODEL_PATH = "unsloth/Meta-Llama-3.1-8B-Instruct"
50
+ BASE_MODEL_PATH = "HiDream-ai/HiDream-I1-Full"
51
+
52
+
53
+ class HidreamModel(BaseModel):
54
+ arch = "hidream"
55
+ hidream_transformer_class = HiDreamImageTransformer2DModel
56
+ hidream_pipeline_class = HiDreamImagePipeline
57
+
58
+ def __init__(
59
+ self,
60
+ device,
61
+ model_config: ModelConfig,
62
+ dtype='bf16',
63
+ custom_pipeline=None,
64
+ noise_scheduler=None,
65
+ **kwargs
66
+ ):
67
+ super().__init__(
68
+ device,
69
+ model_config,
70
+ dtype,
71
+ custom_pipeline,
72
+ noise_scheduler,
73
+ **kwargs
74
+ )
75
+ self.is_flow_matching = True
76
+ self.is_transformer = True
77
+ self.target_lora_modules = ['HiDreamImageTransformer2DModel']
78
+
79
+ # static method to get the noise scheduler
80
+ @staticmethod
81
+ def get_train_scheduler():
82
+ return CustomFlowMatchEulerDiscreteScheduler(**scheduler_config)
83
+
84
+ def get_bucket_divisibility(self):
85
+ return 16
86
+
87
+ def load_model(self):
88
+ dtype = self.torch_dtype
89
+ # HiDream-ai/HiDream-I1-Full
90
+ self.print_and_status_update("Loading HiDream model")
91
+ # will be updated if we detect a existing checkpoint in training folder
92
+ model_path = self.model_config.name_or_path
93
+ extras_path = self.model_config.extras_name_or_path
94
+
95
+ llama_model_path = self.model_config.model_kwargs.get('llama_model_path', LLAMA_MODEL_PATH)
96
+
97
+ scheduler = HidreamModel.get_train_scheduler()
98
+
99
+ self.print_and_status_update("Loading llama 8b model")
100
+
101
+ tokenizer_4 = PreTrainedTokenizerFast.from_pretrained(
102
+ llama_model_path,
103
+ use_fast=False
104
+ )
105
+
106
+ text_encoder_4 = LlamaForCausalLM.from_pretrained(
107
+ llama_model_path,
108
+ output_hidden_states=True,
109
+ output_attentions=True,
110
+ torch_dtype=torch.bfloat16,
111
+ )
112
+ text_encoder_4.to(self.device_torch, dtype=dtype)
113
+
114
+ if self.model_config.quantize_te:
115
+ self.print_and_status_update("Quantizing llama 8b model")
116
+ quantization_type = get_qtype(self.model_config.qtype_te)
117
+ quantize(text_encoder_4, weights=quantization_type)
118
+ freeze(text_encoder_4)
119
+
120
+ if self.low_vram:
121
+ # unload it for now
122
+ text_encoder_4.to('cpu')
123
+
124
+ flush()
125
+
126
+ self.print_and_status_update("Loading transformer")
127
+
128
+ transformer = self.hidream_transformer_class.from_pretrained(
129
+ model_path,
130
+ subfolder="transformer",
131
+ torch_dtype=torch.bfloat16
132
+ )
133
+
134
+ if not self.low_vram:
135
+ transformer.to(self.device_torch, dtype=dtype)
136
+
137
+ if self.model_config.quantize:
138
+ self.print_and_status_update("Quantizing transformer")
139
+ quantization_type = get_qtype(self.model_config.qtype)
140
+ if self.low_vram:
141
+ # move and quantize only certain pieces at a time.
142
+ all_blocks = list(transformer.double_stream_blocks) + list(transformer.single_stream_blocks)
143
+ self.print_and_status_update(" - quantizing transformer blocks")
144
+ for block in tqdm(all_blocks):
145
+ block.to(self.device_torch, dtype=dtype)
146
+ quantize(block, weights=quantization_type)
147
+ freeze(block)
148
+ block.to('cpu')
149
+ # flush()
150
+
151
+ self.print_and_status_update(" - quantizing extras")
152
+ transformer.to(self.device_torch, dtype=dtype)
153
+ quantize(transformer, weights=quantization_type)
154
+ freeze(transformer)
155
+ else:
156
+ quantize(transformer, weights=quantization_type)
157
+ freeze(transformer)
158
+
159
+ if self.low_vram:
160
+ # unload it for now
161
+ transformer.to('cpu')
162
+
163
+ flush()
164
+
165
+ self.print_and_status_update("Loading vae")
166
+
167
+ vae = AutoencoderKL.from_pretrained(
168
+ extras_path,
169
+ subfolder="vae",
170
+ torch_dtype=torch.bfloat16
171
+ ).to(self.device_torch, dtype=dtype)
172
+
173
+
174
+ self.print_and_status_update("Loading clip encoders")
175
+
176
+ text_encoder = CLIPTextModelWithProjection.from_pretrained(
177
+ extras_path,
178
+ subfolder="text_encoder",
179
+ torch_dtype=torch.bfloat16
180
+ ).to(self.device_torch, dtype=dtype)
181
+
182
+ tokenizer = CLIPTokenizer.from_pretrained(
183
+ extras_path,
184
+ subfolder="tokenizer"
185
+ )
186
+
187
+ text_encoder_2 = CLIPTextModelWithProjection.from_pretrained(
188
+ extras_path,
189
+ subfolder="text_encoder_2",
190
+ torch_dtype=torch.bfloat16
191
+ ).to(self.device_torch, dtype=dtype)
192
+
193
+ tokenizer_2 = CLIPTokenizer.from_pretrained(
194
+ extras_path,
195
+ subfolder="tokenizer_2"
196
+ )
197
+
198
+ flush()
199
+ self.print_and_status_update("Loading T5 encoders")
200
+
201
+ text_encoder_3 = T5EncoderModel.from_pretrained(
202
+ extras_path,
203
+ subfolder="text_encoder_3",
204
+ torch_dtype=torch.bfloat16
205
+ ).to(self.device_torch, dtype=dtype)
206
+
207
+ if self.model_config.quantize_te:
208
+ self.print_and_status_update("Quantizing T5")
209
+ quantization_type = get_qtype(self.model_config.qtype_te)
210
+ quantize(text_encoder_3, weights=quantization_type)
211
+ freeze(text_encoder_3)
212
+ flush()
213
+
214
+ tokenizer_3 = T5Tokenizer.from_pretrained(
215
+ extras_path,
216
+ subfolder="tokenizer_3"
217
+ )
218
+ flush()
219
+
220
+ if self.low_vram:
221
+ self.print_and_status_update("Moving everything to device")
222
+ # move it all back
223
+ transformer.to(self.device_torch, dtype=dtype)
224
+ vae.to(self.device_torch, dtype=dtype)
225
+ text_encoder.to(self.device_torch, dtype=dtype)
226
+ text_encoder_2.to(self.device_torch, dtype=dtype)
227
+ text_encoder_4.to(self.device_torch, dtype=dtype)
228
+ text_encoder_3.to(self.device_torch, dtype=dtype)
229
+
230
+ # set to eval mode
231
+ # transformer.eval()
232
+ vae.eval()
233
+ text_encoder.eval()
234
+ text_encoder_2.eval()
235
+ text_encoder_4.eval()
236
+ text_encoder_3.eval()
237
+
238
+ pipe = self.hidream_pipeline_class(
239
+ scheduler=scheduler,
240
+ vae=vae,
241
+ text_encoder=text_encoder,
242
+ tokenizer=tokenizer,
243
+ text_encoder_2=text_encoder_2,
244
+ tokenizer_2=tokenizer_2,
245
+ text_encoder_3=text_encoder_3,
246
+ tokenizer_3=tokenizer_3,
247
+ text_encoder_4=text_encoder_4,
248
+ tokenizer_4=tokenizer_4,
249
+ transformer=transformer,
250
+ )
251
+
252
+ flush()
253
+
254
+ text_encoder_list = [text_encoder, text_encoder_2, text_encoder_3, text_encoder_4]
255
+ tokenizer_list = [tokenizer, tokenizer_2, tokenizer_3, tokenizer_4]
256
+
257
+ for te in text_encoder_list:
258
+ # set the dtype
259
+ te.to(self.device_torch, dtype=dtype)
260
+ # freeze the model
261
+ freeze(te)
262
+ # set to eval mode
263
+ te.eval()
264
+ # set the requires grad to false
265
+ te.requires_grad_(False)
266
+
267
+ flush()
268
+
269
+ # save it to the model class
270
+ self.vae = vae
271
+ self.text_encoder = text_encoder_list # list of text encoders
272
+ self.tokenizer = tokenizer_list # list of tokenizers
273
+ self.model = pipe.transformer
274
+ self.pipeline = pipe
275
+ self.print_and_status_update("Model Loaded")
276
+
277
+ def get_generation_pipeline(self):
278
+ scheduler = FlowUniPCMultistepScheduler(
279
+ num_train_timesteps=1000,
280
+ shift=3.0,
281
+ use_dynamic_shifting=False
282
+ )
283
+
284
+ pipeline: HiDreamImagePipeline = HiDreamImagePipeline(
285
+ scheduler=scheduler,
286
+ vae=self.vae,
287
+ text_encoder=self.text_encoder[0],
288
+ tokenizer=self.tokenizer[0],
289
+ text_encoder_2=self.text_encoder[1],
290
+ tokenizer_2=self.tokenizer[1],
291
+ text_encoder_3=self.text_encoder[2],
292
+ tokenizer_3=self.tokenizer[2],
293
+ text_encoder_4=self.text_encoder[3],
294
+ tokenizer_4=self.tokenizer[3],
295
+ transformer=unwrap_model(self.model),
296
+ aggressive_unloading=self.low_vram
297
+ )
298
+
299
+ pipeline = pipeline.to(self.device_torch)
300
+
301
+ return pipeline
302
+
303
+ def generate_single_image(
304
+ self,
305
+ pipeline: HiDreamImagePipeline,
306
+ gen_config: GenerateImageConfig,
307
+ conditional_embeds: PromptEmbeds,
308
+ unconditional_embeds: PromptEmbeds,
309
+ generator: torch.Generator,
310
+ extra: dict,
311
+ ):
312
+ img = pipeline(
313
+ prompt_embeds=conditional_embeds.text_embeds,
314
+ pooled_prompt_embeds=conditional_embeds.pooled_embeds,
315
+ negative_prompt_embeds=unconditional_embeds.text_embeds,
316
+ negative_pooled_prompt_embeds=unconditional_embeds.pooled_embeds,
317
+ height=gen_config.height,
318
+ width=gen_config.width,
319
+ num_inference_steps=gen_config.num_inference_steps,
320
+ guidance_scale=gen_config.guidance_scale,
321
+ latents=gen_config.latents,
322
+ generator=generator,
323
+ **extra
324
+ ).images[0]
325
+ return img
326
+
327
+ def get_noise_prediction(
328
+ self,
329
+ latent_model_input: torch.Tensor,
330
+ timestep: torch.Tensor, # 0 to 1000 scale
331
+ text_embeddings: PromptEmbeds,
332
+ **kwargs
333
+ ):
334
+ batch_size = latent_model_input.shape[0]
335
+ with torch.no_grad():
336
+ if latent_model_input.shape[-2] != latent_model_input.shape[-1]:
337
+ B, C, H, W = latent_model_input.shape
338
+ pH, pW = H // self.model.config.patch_size, W // self.model.config.patch_size
339
+
340
+ img_sizes = torch.tensor([pH, pW], dtype=torch.int64).reshape(-1)
341
+ img_ids = torch.zeros(pH, pW, 3)
342
+ img_ids[..., 1] = img_ids[..., 1] + torch.arange(pH)[:, None]
343
+ img_ids[..., 2] = img_ids[..., 2] + torch.arange(pW)[None, :]
344
+ img_ids = img_ids.reshape(pH * pW, -1)
345
+ img_ids_pad = torch.zeros(self.transformer.max_seq, 3)
346
+ img_ids_pad[:pH*pW, :] = img_ids
347
+
348
+ img_sizes = img_sizes.unsqueeze(0).to(latent_model_input.device)
349
+ img_sizes = torch.cat([img_sizes] * batch_size, dim=0)
350
+ img_ids = img_ids_pad.unsqueeze(0).to(latent_model_input.device)
351
+ img_ids = torch.cat([img_ids] * batch_size, dim=0)
352
+ else:
353
+ img_sizes = img_ids = None
354
+
355
+ dtype = self.model.dtype
356
+ device = self.device_torch
357
+
358
+ # Pack the latent
359
+ if latent_model_input.shape[-2] != latent_model_input.shape[-1]:
360
+ B, C, H, W = latent_model_input.shape
361
+ patch_size = self.transformer.config.patch_size
362
+ pH, pW = H // patch_size, W // patch_size
363
+ out = torch.zeros(
364
+ (B, C, self.transformer.max_seq, patch_size * patch_size),
365
+ dtype=latent_model_input.dtype,
366
+ device=latent_model_input.device
367
+ )
368
+ latent_model_input = einops.rearrange(latent_model_input, 'B C (H p1) (W p2) -> B C (H W) (p1 p2)', p1=patch_size, p2=patch_size)
369
+ out[:, :, 0:pH*pW] = latent_model_input
370
+ latent_model_input = out
371
+
372
+ text_embeds = text_embeddings.text_embeds
373
+ # run the to for the list
374
+ text_embeds = [te.to(device, dtype=dtype) for te in text_embeds]
375
+
376
+ noise_pred = self.transformer(
377
+ hidden_states = latent_model_input,
378
+ timesteps = timestep,
379
+ encoder_hidden_states = text_embeds,
380
+ pooled_embeds = text_embeddings.pooled_embeds.to(device, dtype=dtype),
381
+ img_sizes = img_sizes,
382
+ img_ids = img_ids,
383
+ return_dict = False,
384
+ )[0]
385
+ noise_pred = -noise_pred
386
+
387
+ return noise_pred
388
+
389
+ def get_prompt_embeds(self, prompt: str) -> PromptEmbeds:
390
+ self.text_encoder_to(self.device_torch, dtype=self.torch_dtype)
391
+ max_sequence_length = 128
392
+ prompt_embeds, pooled_prompt_embeds = self.pipeline._encode_prompt(
393
+ prompt = prompt,
394
+ prompt_2 = prompt,
395
+ prompt_3 = prompt,
396
+ prompt_4 = prompt,
397
+ device = self.device_torch,
398
+ dtype = self.torch_dtype,
399
+ num_images_per_prompt = 1,
400
+ max_sequence_length = max_sequence_length,
401
+ )
402
+ pe = PromptEmbeds(
403
+ [prompt_embeds, pooled_prompt_embeds]
404
+ )
405
+ return pe
406
+
407
+ def get_model_has_grad(self):
408
+ # return from a weight if it has grad
409
+ return self.model.double_stream_blocks[0].block.attn1.to_q.weight.requires_grad
410
+
411
+ def get_te_has_grad(self):
412
+ # assume no one wants to finetune 4 text encoders.
413
+ return False
414
+
415
+ def save_model(self, output_path, meta, save_dtype):
416
+ # only save the unet
417
+ transformer: HiDreamImageTransformer2DModel = unwrap_model(self.model)
418
+ transformer.save_pretrained(
419
+ save_directory=os.path.join(output_path, 'transformer'),
420
+ safe_serialization=True,
421
+ )
422
+
423
+ meta_path = os.path.join(output_path, 'aitk_meta.yaml')
424
+ with open(meta_path, 'w') as f:
425
+ yaml.dump(meta, f)
426
+
427
+ def get_loss_target(self, *args, **kwargs):
428
+ noise = kwargs.get('noise')
429
+ batch = kwargs.get('batch')
430
+ return (noise - batch.latents).detach()
431
+
432
+ def get_transformer_block_names(self) -> Optional[List[str]]:
433
+ return ['double_stream_blocks', 'single_stream_blocks']
434
+
435
+ def convert_lora_weights_before_save(self, state_dict):
436
+ # currently starte with transformer. but needs to start with diffusion_model. for comfyui
437
+ new_sd = {}
438
+ for key, value in state_dict.items():
439
+ new_key = key.replace("transformer.", "diffusion_model.")
440
+ new_sd[new_key] = value
441
+ return new_sd
442
+
443
+ def convert_lora_weights_before_load(self, state_dict):
444
+ # saved as diffusion_model. but needs to be transformer. for ai-toolkit
445
+ new_sd = {}
446
+ for key, value in state_dict.items():
447
+ new_key = key.replace("diffusion_model.", "transformer.")
448
+ new_sd[new_key] = value
449
+ return new_sd
450
+
451
+ def get_base_model_version(self):
452
+ return "hidream_i1"
453
+