Lora Loading Log and Z Image Distill (#450)

This commit is contained in:
Bubbliiiing
2026-02-04 16:55:24 +08:00
committed by GitHub
parent a6b026526f
commit ee44fbc950
6 changed files with 3859 additions and 7 deletions
File diff suppressed because it is too large Load Diff
+34
View File
@@ -0,0 +1,34 @@
export MODEL_NAME="models/Diffusion_Transformer/Z-Image"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/z_image/train_distill.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--train_batch_size=1 \
--image_sample_size=1328 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=100 \
--learning_rate=2e-05 \
--learning_rate_critic=2e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_z_image_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--trainable_modules "." \
--random_hw_adapt \
--uniform_sampling
File diff suppressed because it is too large Load Diff
+35
View File
@@ -0,0 +1,35 @@
export MODEL_NAME="models/Diffusion_Transformer/Z-Image"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/z_image/train_distill_lora.py \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--train_batch_size=1 \
--image_sample_size=1328 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--learning_rate_critic=1e-05 \
--seed=42 \
--output_dir="output_dir_z_image_distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--rank=128 \
--network_alpha=128 \
--target_name="to_q,to_k,to_v,feed_forward.w1,feed_forward.w2,feed_forward.w3" \
--random_hw_adapt \
--use_peft_lora \
--uniform_sampling
+1 -1
View File
@@ -28,6 +28,6 @@ accelerate launch --mixed_precision="bf16" scripts/z_image/train_lora.py \
--enable_bucket \
--rank=64 \
--network_alpha=64 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
--target_name="to_q,to_k,to_v,feed_forward.w1,feed_forward.w2,feed_forward.w3" \
--use_peft_lora \
--uniform_sampling
+55 -6
View File
@@ -402,12 +402,23 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
if lora_path is None:
return pipeline
print(f"[LoRA Merge] Starting to merge LoRA from: {lora_path if lora_path else 'state_dict'}")
print(f"[LoRA Merge] Multiplier: {multiplier}, Device: {device}, Dtype: {dtype}")
LORA_PREFIX_TRANSFORMER = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
if state_dict is None:
state_dict = load_file(lora_path)
if lora_path.endswith("safetensors"):
print(f"[LoRA Merge] Loading safetensors file...")
state_dict = load_file(lora_path)
else:
print(f"[LoRA Merge] Loading pytorch file...")
state_dict = torch.load(lora_path, map_location="cpu")
else:
print(f"[LoRA Merge] Using provided state_dict")
state_dict = state_dict
updates = defaultdict(dict)
for key, value in state_dict.items():
if "diffusion_model." in key:
@@ -436,16 +447,24 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
layer, elem = key.split('.', 1)
updates[layer][elem] = value
print(f"[LoRA Merge] Organized into {len(updates)} layers")
sequential_cpu_offload_flag = False
if pipeline.transformer.device == torch.device(type="meta"):
print(f"[LoRA Merge] Removing hooks for meta device...")
pipeline.remove_all_hooks()
sequential_cpu_offload_flag = True
offload_device = pipeline._offload_device
merged_count = 0
skipped_count = 0
error_count = 0
for layer, elems in updates.items():
if "lora_te" in layer:
if transformer_only:
skipped_count += 1
continue
else:
layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
@@ -472,7 +491,7 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
break
except Exception:
if len(layer_infos) == 0:
print(f'Error loading layer in front search: {layer}. Try it in back search.')
print(f'[LoRA Merge] Warning: Error loading layer in front search: {layer}. Try it in back search.')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
@@ -480,6 +499,7 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
except Exception:
if "lora_te" in layer:
if transformer_only:
skipped_count += 1
continue
else:
layer_infos = layer.split(LORA_PREFIX_TEXT_ENCODER + "_")[-1].split("_")
@@ -496,7 +516,7 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
while start_index < len_layer_infos:
try:
if start_index >= end_indx:
print(f'Error loading layer in back search: {layer}')
print(f'[LoRA Merge] Error: Failed to load layer in back search: {layer}')
error_flag = True
break
curr_layer = curr_layer.__getattr__("_".join(layer_infos[start_index:end_indx]))
@@ -505,6 +525,7 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
except Exception:
end_indx -= 1
if error_flag:
error_count += 1
continue
origin_dtype = curr_layer.weight.data.dtype
@@ -526,9 +547,15 @@ def merge_lora(pipeline, lora_path, multiplier, device='cpu', dtype=torch.float3
else:
curr_layer.weight.data += multiplier * alpha * torch.mm(weight_up, weight_down)
curr_layer = curr_layer.to(origin_device, origin_dtype)
merged_count += 1
print(f"[LoRA Merge] Completed: {merged_count} layers merged, {skipped_count} layers skipped, {error_count} errors")
if sequential_cpu_offload_flag:
print(f"[LoRA Merge] Re-enabling sequential CPU offload...")
pipeline.enable_sequential_cpu_offload(device=offload_device)
print(f"[LoRA Merge] ✓ LoRA merge finished successfully")
return pipeline
# TODO: Refactor with merge_lora.
@@ -536,10 +563,19 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
if lora_path is None:
return pipeline
print(f"[LoRA Unmerge] Starting to unmerge LoRA from: {lora_path}")
print(f"[LoRA Unmerge] Multiplier: {multiplier}, Device: {device}, Dtype: {dtype}")
"""Unmerge state_dict in LoRANetwork from the pipeline in diffusers."""
LORA_PREFIX_UNET = "lora_unet"
LORA_PREFIX_TEXT_ENCODER = "lora_te"
state_dict = load_file(lora_path)
if lora_path.endswith("safetensors"):
print(f"[LoRA Unmerge] Loading safetensors file...")
state_dict = load_file(lora_path)
else:
print(f"[LoRA Unmerge] Loading pytorch file...")
state_dict = torch.load(lora_path, map_location="cpu")
updates = defaultdict(dict)
for key, value in state_dict.items():
@@ -569,11 +605,17 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
layer, elem = key.split('.', 1)
updates[layer][elem] = value
print(f"[LoRA Unmerge] Organized into {len(updates)} layers")
sequential_cpu_offload_flag = False
if pipeline.transformer.device == torch.device(type="meta"):
print(f"[LoRA Unmerge] Removing hooks for meta device...")
pipeline.remove_all_hooks()
sequential_cpu_offload_flag = True
unmerged_count = 0
error_count = 0
for layer, elems in updates.items():
if "lora_te" in layer:
@@ -601,7 +643,7 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
break
except Exception:
if len(layer_infos) == 0:
print(f'Error loading layer in front search: {layer}. Try it in back search.')
print(f'[LoRA Unmerge] Warning: Error loading layer in front search: {layer}. Try it in back search.')
if len(temp_name) > 0:
temp_name += "_" + layer_infos.pop(0)
else:
@@ -622,7 +664,7 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
while start_index < len_layer_infos:
try:
if start_index >= end_indx:
print(f'Error loading layer in back search: {layer}')
print(f'[LoRA Unmerge] Error: Failed to load layer in back search: {layer}')
error_flag = True
break
curr_layer = curr_layer.__getattr__("_".join(layer_infos[start_index:end_indx]))
@@ -631,6 +673,7 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
except Exception:
end_indx -= 1
if error_flag:
error_count += 1
continue
origin_dtype = curr_layer.weight.data.dtype
@@ -652,7 +695,13 @@ def unmerge_lora(pipeline, lora_path, multiplier=1, device="cpu", dtype=torch.fl
else:
curr_layer.weight.data -= multiplier * alpha * torch.mm(weight_up, weight_down)
curr_layer = curr_layer.to(origin_device, origin_dtype)
unmerged_count += 1
print(f"[LoRA Unmerge] Completed: {unmerged_count} layers unmerged, {error_count} errors")
if sequential_cpu_offload_flag:
print(f"[LoRA Unmerge] Re-enabling sequential CPU offload...")
pipeline.enable_sequential_cpu_offload(device=device)
print(f"[LoRA Unmerge] ✓ LoRA unmerge finished successfully")
return pipeline