| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
|
|
| from accelerate import DistributedType |
| from diffusers import SanaPipeline, SanaTransformer2DModel |
| from diffusers.training_utils import cast_training_params |
| from peft.utils import get_peft_model_state_dict |
|
|
|
|
| |
| def create_save_model_hook( |
| accelerator, |
| unwrap_model, |
| transformer, |
| ): |
| def save_model_hook(models, weights, output_dir): |
| if accelerator.is_main_process: |
| transformer_lora_layers_to_save = None |
|
|
| for model in models: |
| if isinstance(unwrap_model(model), type(unwrap_model(transformer))): |
| transformer_model = unwrap_model(model) |
| transformer_lora_layers_to_save = get_peft_model_state_dict( |
| transformer_model |
| ) |
| else: |
| raise ValueError(f"unexpected save model: {model.__class__}") |
|
|
| |
| if weights: |
| weights.pop() |
|
|
| SanaPipeline.save_lora_weights( |
| output_dir, |
| transformer_lora_layers=transformer_lora_layers_to_save, |
| ) |
|
|
| return save_model_hook |
|
|
|
|
| |
| def create_load_model_hook( |
| accelerator, |
| unwrap_model, |
| transformer, |
| args, |
| ): |
| def load_model_hook(models, output_dir): |
| transformer_ = None |
|
|
| if not accelerator.distributed_type == DistributedType.DEEPSPEED: |
| while len(models) > 0: |
| model = models.pop() |
|
|
| if isinstance(unwrap_model(model), type(unwrap_model(transformer))): |
| transformer_ = model |
| else: |
| raise ValueError(f"unexpected save model: {model.__class__}") |
| else: |
| transformer_ = SanaTransformer2DModel.from_pretrained( |
| args.pretrained_model_name_or_path, |
| subfolder="transformer", |
| local_files_only=True, |
| ) |
|
|
| |
| |
| |
| if args.mixed_precision == "fp16": |
| models = [transformer_] |
| |
| cast_training_params(models) |
|
|
| return load_model_hook |
|
|