Skip to content

Commit 024a29f

Browse files
shenliao.slaPanAndy
authored andcommitted
(feat): add merge lora scripts and update Reward-FL docs.
1 parent 74f8da6 commit 024a29f

2 files changed

Lines changed: 117 additions & 0 deletions

File tree

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,94 @@
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+
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
import json
2+
from safetensors import safe_open
3+
from safetensors.torch import save_file
4+
import torch
5+
6+
shard_files = ['diffusion_pytorch_model-00001-of-00006.safetensors', 'diffusion_pytorch_model-00002-of-00006.safetensors',
7+
"diffusion_pytorch_model-00003-of-00006.safetensors","diffusion_pytorch_model-00004-of-00006.safetensors",
8+
"diffusion_pytorch_model-00005-of-00006.safetensors","diffusion_pytorch_model-00006-of-00006.safetensors"]
9+
print("Shard files to load:", shard_files)
10+
11+
full_state_dict = {}
12+
13+
for shard_file in shard_files:
14+
print(f"Loading {shard_file}...")
15+
with safe_open(shard_file, framework="pt", device="cpu") as f:
16+
for key in f.keys():
17+
full_state_dict[key] = f.get_tensor(key).to(dtype=torch.bfloat16)
18+
19+
output_file = "diffusion_pytorch_model.safetensors"
20+
print(f"Saving merged model to {output_file}...")
21+
save_file(full_state_dict, output_file)
22+
23+
print("Done! Merged model saved.")

0 commit comments

Comments
 (0)