|
| 1 | +import os |
| 2 | +import argparse |
| 3 | +import torch |
| 4 | +from safetensors import safe_open |
| 5 | +from safetensors.torch import save_file |
| 6 | + |
| 7 | +def load_state_dict_from_safetensors(file_path, device="cpu"): |
| 8 | + state_dict = {} |
| 9 | + with safe_open(file_path, framework="pt", device=str(device)) as f: |
| 10 | + for k in f.keys(): |
| 11 | + state_dict[k] = f.get_tensor(k) |
| 12 | + return state_dict |
| 13 | + |
| 14 | +def merge_lora_into_state_dict( |
| 15 | + base_state_dict, |
| 16 | + lora_state_dict, |
| 17 | + alpha=1.0, |
| 18 | + device="cpu", |
| 19 | + dtype=torch.bfloat16 |
| 20 | +): |
| 21 | + lora_name_map = {} |
| 22 | + for key in lora_state_dict: |
| 23 | + if ".lora_B." not in key: |
| 24 | + continue |
| 25 | + clean_key = key.replace("base_model.model.", "") |
| 26 | + parts = clean_key.split(".") |
| 27 | + try: |
| 28 | + lora_b_idx = parts.index("lora_B") |
| 29 | + except ValueError: |
| 30 | + continue |
| 31 | + target_parts = parts[:lora_b_idx] + parts[lora_b_idx + 2:] # 跳过 lora_B 和 rank idx |
| 32 | + if target_parts[0] == "diffusion_model": |
| 33 | + target_parts = target_parts[1:] |
| 34 | + target_name = ".".join(target_parts) |
| 35 | + lora_A_key = key.replace(".lora_B.", ".lora_A.").replace("base_model.model.", "") |
| 36 | + lora_name_map[target_name] = (key, lora_A_key) |
| 37 | + |
| 38 | + merged_state_dict = base_state_dict.copy() |
| 39 | + updated = 0 |
| 40 | + |
| 41 | + for target_name, (lora_B_key, lora_A_key) in lora_name_map.items(): |
| 42 | + if target_name not in merged_state_dict: |
| 43 | + print(f"Warning: {target_name} not in base model. Skipping.") |
| 44 | + continue |
| 45 | + |
| 46 | + weight_orig = merged_state_dict[target_name].to(device=device, dtype=dtype) |
| 47 | + weight_up = lora_state_dict[lora_B_key].to(device=device, dtype=dtype) |
| 48 | + weight_down = lora_state_dict[lora_A_key].to(device=device, dtype=dtype) |
| 49 | + |
| 50 | + if len(weight_up.shape) == 4: |
| 51 | + weight_up = weight_up.squeeze(-1).squeeze(-1) # [out, r, 1, 1] -> [out, r] |
| 52 | + weight_down = weight_down.squeeze(-1).squeeze(-1) # [r, in, 1, 1] -> [r, in] |
| 53 | + lora_weight = alpha * (weight_up @ weight_down).unsqueeze(-1).unsqueeze(-1) |
| 54 | + else: |
| 55 | + lora_weight = alpha * (weight_up @ weight_down) |
| 56 | + |
| 57 | + merged_state_dict[target_name] = weight_orig + lora_weight |
| 58 | + updated += 1 |
| 59 | + |
| 60 | + print(f"Merged {updated} LoRA adapters into base model.") |
| 61 | + return merged_state_dict |
| 62 | + |
| 63 | +parser = argparse.ArgumentParser() |
| 64 | +parser.add_argument('--lora_dir', type=str, default="/data/models/Wan2.2-I2V-A14B-4steps-lora-rank64-Seko-V1/high_noise_model.safetensors") |
| 65 | +parser.add_argument('--base_model_path', type=str, default="/data/models/Wan22_base/high_noise_model/diffusion_pytorch_model.safetensors") |
| 66 | +parser.add_argument('--save_dir', type=str, default="/data/models/Wan22/high_noise_model/diffusion_pytorch_model.safetensors") |
| 67 | +parser.add_argument('--alpha', type=float, default=1.0) |
| 68 | +args = parser.parse_args() |
| 69 | + |
| 70 | +print("Loading base model state dict...") |
| 71 | +base_sd = load_state_dict_from_safetensors(args.base_model_path) |
| 72 | + |
| 73 | +# load lora |
| 74 | +lora_sd = load_state_dict_from_safetensors(args.lora_dir) |
| 75 | + |
| 76 | +clean_lora_sd = {} |
| 77 | +for k, v in lora_sd.items(): |
| 78 | + clean_k = k.replace("base_model.model.", "") |
| 79 | + clean_lora_sd[clean_k] = v |
| 80 | + |
| 81 | +# merge |
| 82 | +merged_sd = merge_lora_into_state_dict( |
| 83 | + base_state_dict=base_sd, |
| 84 | + lora_state_dict=clean_lora_sd, |
| 85 | + alpha=args.alpha, |
| 86 | + device="cpu", |
| 87 | + dtype=torch.bfloat16 |
| 88 | +) |
| 89 | + |
| 90 | +# save |
| 91 | +os.makedirs(os.path.dirname(args.save_dir), exist_ok=True) |
| 92 | +save_file(merged_sd, args.save_dir) |
| 93 | +print(f"Merged model saved to {args.save_dir}") |
| 94 | + |
0 commit comments