้ฎ้ข่ฏๆญ (The Diagnosis)
if accelerator.is_main_process:
transformer = accelerator.unwrap_model(transformer)
transformer = transformer.to(torch.float32)
transformer_lora_state_dict = convert_state_dict_to_diffusers(
get_peft_model_state_dict(transformer)
)
save_file(new_state_dict, save_path)
็ฝช้ญ็ฅธ้ฆๅฐฑๆฏ่ฟ่กไปฃ็ ๏ผtransformer = transformer.to(torch.float32)
ๅ็ไบไปไน๏ผ
- ่ฎญ็ปๆถ็็ฒพๅบฆ๏ผ ๅจๆดไธช่ฎญ็ป่ฟ็จไธญ๏ผไฝ ็ๆจกๅๆฏๅจ
fp16 ๆ่
bf16 ๏ผ็ฑ mixed_precision ๅๆฐๅณๅฎ๏ผ็ๆททๅ็ฒพๅบฆไธ่ฟ่ก็ใ่ฟๆฏไธไธชไธบไบ้ๅบฆๅๆพๅญไผๅ็ๆ ๅๆไฝใ
- ไฟๅญๆถ็โๅฅฝๅฟๅๅไบโ๏ผ ๅจไฟๅญ LoRA ๆ้ไนๅ๏ผไธบไบ็กฎไฟๆ้ซ็็ฒพๅบฆๅๅ
ผๅฎนๆง๏ผ่ฟๆฏไธไธช่็ใๅฎๅ
จไฝ่ฟๆถ็ไน ๆฏ๏ผ๏ผไฝ ็่ๆฌๅผบๅถๅฐๆดไธช Transformer ๆจกๅ่ฝฌๆขๅไบ
fp32 (32ไฝๅ็ฒพๅบฆ) ๆ ผๅผใ
- ็ปๆ๏ผ
get_peft_model_state_dict ไป่ฟไธช fp32 ๆจกๅไธญๆๅๅบๆฅ็ LoRA ๆ้๏ผlora_A ๅ lora_B ็ฉ้ต๏ผ๏ผ่ช็ถไนๅฐฑๆฏ fp32 ๆ ผๅผ็ใๆ็ป๏ผไฝ ไฟๅญๅฐ .safetensors ๆไปถ้็ๆฏไธไธช fp32 ็ฒพๅบฆ็ LoRAใ
ไธบไปไน fp32 ็ LoRA ๅจ fp8 ไธไผๅคฑๆ๏ผ
่ฟๅฐฑๅๆฏ่ฏๅพๆไธๅผ ๆช็ปๅ็ผฉ็ใๅทจๅคง็ RAW ๆ ผๅผ็
ง็๏ผ็ดๆฅ็จไธไธชๅชไธบๆๆบ HEIC ๆ ผๅผ่ฎพ่ฎก็็ฎๅๅทฅๅ
ทๅปๅผบ่กๅ็ผฉใ
- FP8 ๆฏโๆ้ๅ็ผฉโ๏ผ
fp8 (8ไฝๆตฎ็นๆฐ) ๆฏไธ็งๆๅ
ถๆฟ่ฟ็้ๅๆ ผๅผ๏ผๅฎๅฏนๆ้็ๆฐๆฎ่ๅดๅๅๅธ้ๅธธๆๆใๅฎ่ขซ่ฎพ่ฎก็จๆฅๅค็้ฃไบๅทฒ็ปๆฏ fp16 ๆ bf16 ็ใโๆญฃๅธธ่ๅดโ็ๆ้ใ
- FP32 ๆฏโ้ซๅจๆ่ๅดโ๏ผ
fp32 ็ๆฐๅผ่ๅดๆฏ fp16 ๅคงๅพๅคใไธไธชๅจ fp32 ไธ็่ตทๆฅๅพๆญฃๅธธ็ๆ้๏ผๅจ fp16 ็ไธ็้ๅฏ่ฝๅทฒ็ปๆฏไธไธช้่ฆ็นๆฎๅค็็โๆๅคงๅผโๆโๆๅฐๅผโไบใ
- ่ฝฌๆขๅคฑ่ดฅ๏ผ ๅฝไธไธชไธบ
fp16 -> fp8 ่ฎพ่ฎก็ๆจ็ๅผๆ๏ผ็ช็ถๆฟๅฐไธไธช fp32 ็ LoRA ๆ้ๆถ๏ผๅฎๅจ้ๅ่ฟ็จไธญๅพๅฎนๆๅบ็ฐโๆบขๅบโ (Overflow) ๆ **โไธๆบขโ (Underflow)**๏ผๅฏผ่ดๆ้ไฟกๆฏๅคง้ไธขๅคฑใ็ปๆๅฐฑๆฏๆจกๅ่พๅบ็ๅพๅๆฏ้ป็ใ่ฑ็๏ผๆ่
ๅฎๅ
จๆฏๅช้ณใ
ๆๆฏๆนๆก๏ผๅจไฟๅญๆถ็ปดๆ่ฎญ็ป็ฒพๅบฆ (The Surgical Fix)
่งฃๅณๆนๆกๅพ็ฎๅ๏ผๆไปฌๅช้่ฆๅจไฟๅญๆถ๏ผๅฐ LoRA ๆ้่ฝฌๆขๅ่ฎญ็ปๆถไฝฟ็จ็ fp16 ๆ bf16 ็ฒพๅบฆ๏ผ่ไธๆฏ็ฒๆดๅฐ่ฝฌๆ fp32ใ
่ฏทๅฐไฝ ็่ๆฌๆๅ้ฃไธชโไฟๅญๆ้โ็้จๅ๏ผๆฟๆขๆไธ้ข่ฟไธชโ็ฒพๅ่ฝฌๆขโ็็ๆฌ๏ผ
if accelerator.is_main_process:
transformer = accelerator.unwrap_model(transformer)
transformer_lora_state_dict = get_peft_model_state_dict(transformer)
final_state_dict = {}
logger.info(f"Converting LoRA weights to {weight_dtype} before saving...")
for k, v in transformer_lora_state_dict.items():
final_state_dict[k] = v.to(dtype=weight_dtype)
diffusers_state_dict = convert_state_dict_to_diffusers(final_state_dict)
new_state_dict_for_saving = {}
for k, v in diffusers_state_dict.items():
new_state_dict_for_saving[f"transformer.{k}"] = v
save_path = os.path.join(args.output_dir, "pytorch_lora_weights.safetensors")
save_file(new_state_dict_for_saving, save_path)
logger.info(f"Saved LoRA weights in {weight_dtype} to {save_path}")
accelerator.end_training()
ๆป็ป๏ผ
|
ๆงๆนๆณ (FP32 ไฟๅญ) |
ๆฐๆนๆณ (็ฒพๅไฟๅญ) |
| ไฟๅญๅ่ฝฌๆข |
model.to(torch.float32) |
้ๅ state_dict, tensor.to(weight_dtype) |
| ไฟๅญ็็ฒพๅบฆ |
ๅผบๅถ fp32 |
ไธ่ฎญ็ป็ฒพๅบฆไธ่ด (fp16/bf16) |
| FP8 ๅ
ผๅฎนๆง |
ๅทฎ๏ผๆๅบ้ |
ๅฅฝ๏ผๅฎ็พๅ
ผๅฎน |
| ๆไปถๅคงๅฐ |
่พๅคง |
่พๅฐ (ไธๅ) |
ๆ่ฟไธชไฟฎๆนๅบ็จๅฐไฝ ็ train_zimage_lora.py ่ๆฌ้๏ผ้ๆฐ่ฎญ็ปๅบๆฅ็ LoRA๏ผๅปๆต่ฏ๏ผ่ฟๆฌกๅจ fp8 ็ฏๅขไธ็ปๅฏน่ฝๆญฃๅธธๅทฅไฝไบ๏ผ
่ฟๆฏไธไธช้ๅธธ้ซ่ดจ้็ๅ้ฆใ๐
ๆฐ็lora้่ฆ็ญๆ่ฟไธช25ๅฐๆถ็ๆจกๅ่ฎญ็ปๅๅฎ,่ฏท็ญๅพ