Skip to content

Commit 65b506c

Browse files
mrloopjianghaicheng
authored andcommitted
[feat] support wan2.1 fsdp + cp
Change-Id: I6732bfb5ff0e962515330f5a07ff18d79664cac9
1 parent ad97cc9 commit 65b506c

22 files changed

Lines changed: 773 additions & 130 deletions

configs/models/wan/wan2_1_i2v.yaml

Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
_target_: loongforge.models.diffusion.WanConfig
2+
num_layers: 40
3+
hidden_size: 5120
4+
ffn_hidden_size: 13824
5+
text_dim: 4096
6+
num_attention_heads: 40
7+
latent_in_channels: 4
8+
latent_out_channels: 8
9+
latent_patch_size: [1, 2, 2]
10+
latent_space_scale: 0.5
11+
latent_time_scale: 1.0
12+
vae_temporal_compress: 4
13+
vae_spatial_compress: 8
14+
has_image_input: True
15+
clip_num_image_tokens: 257
16+
freq_dim: 256
17+
in_dim: 36
18+
out_dim: 16
19+
norm_epsilon: 1e-6
20+
require_clip_embedding: True
21+
has_image_pos_emb: False
22+
model_type: "wan2_1_i2v"
23+
attention_dropout: 0.0
24+
hidden_dropout: 0.0
25+
use_fused_wan_rope: True

configs/models/wan/wan2_2_i2v.yaml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,3 +19,6 @@ norm_epsilon: 1e-6
1919
require_clip_embedding: False
2020
has_image_pos_emb: False
2121
model_type: "wan2_2_i2v"
22+
use_fused_wan_rope: True
23+
attention_dropout: 0.0
24+
hidden_dropout: 0.0

examples/wan/convert_checkpoint_hg2mcore.py

Lines changed: 39 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -117,6 +117,16 @@ def load_huggingface_chekckpoints(path, num_checkpoints):
117117
"time_projection.1.bias",
118118
]
119119
extra_first_part_dict = {
120+
"wan2_1_i2v": [
121+
"img_emb.proj.0.weight",
122+
"img_emb.proj.0.bias",
123+
"img_emb.proj.1.weight",
124+
"img_emb.proj.1.bias",
125+
"img_emb.proj.3.weight",
126+
"img_emb.proj.3.bias",
127+
"img_emb.proj.4.weight",
128+
"img_emb.proj.4.bias",
129+
],
120130
"wan2_2_i2v": []
121131
}
122132
first_part_list = base_first_part_list + extra_first_part_dict.get(model_name, [])
@@ -158,6 +168,15 @@ def load_huggingface_chekckpoints(path, num_checkpoints):
158168
"decoder.layers.0.cross_attn.q_layernorm.weight": "blocks.0.cross_attn.norm_q.weight",
159169
"decoder.layers.0.cross_attn.k_layernorm.weight": "blocks.0.cross_attn.norm_k.weight",
160170
}
171+
wan2_1_second_part_dict = {
172+
"decoder.layers.0.cross_attn.linear_k_img.weight": "blocks.0.cross_attn.k_img.weight",
173+
"decoder.layers.0.cross_attn.linear_k_img.bias": "blocks.0.cross_attn.k_img.bias",
174+
"decoder.layers.0.cross_attn.linear_v_img.weight": "blocks.0.cross_attn.v_img.weight",
175+
"decoder.layers.0.cross_attn.linear_v_img.bias": "blocks.0.cross_attn.v_img.bias",
176+
"decoder.layers.0.cross_attn.k_img_layernorm.weight": "blocks.0.cross_attn.norm_k_img.weight",
177+
}
178+
if model_name == "wan2_1_i2v":
179+
second_part_dict.update(wan2_1_second_part_dict)
161180
# Parts that do not need transpose inside
162181
inside_blk_replace_dict = {
163182
"blocks.0.ffn.0.weight": "decoder.layers.0.ffn.0.weight",
@@ -171,6 +190,11 @@ def load_huggingface_chekckpoints(path, num_checkpoints):
171190
"blocks.0.cross_attn.norm_q.weight": "decoder.layers.0.cross_attn.q_layernorm.weight",
172191
"blocks.0.cross_attn.norm_k.weight": "decoder.layers.0.cross_attn.k_layernorm.weight",
173192
}
193+
wan2_1_inside_blk_replace_dict = {
194+
"blocks.0.cross_attn.norm_k_img.weight": "decoder.layers.0.cross_attn.k_img_layernorm.weight",
195+
}
196+
if model_name == "wan2_1_i2v":
197+
inside_blk_replace_dict.update(wan2_1_inside_blk_replace_dict)
174198
# Model last part
175199
third_part_dict = {
176200
"head.modulation",
@@ -243,6 +267,21 @@ def load_huggingface_chekckpoints(path, num_checkpoints):
243267
concat_kv_bias
244268
)
245269

270+
if model_name == "wan2_1_i2v":
271+
cross_k_img_w = state_dict["blocks." + str(i) + ".cross_attn.k_img.weight"]
272+
cross_k_img_w = rearrange(cross_k_img_w, "(R N D) H -> (N R D) H", R=1, N=40, D=128, H=5120)
273+
cross_k_img_b = state_dict["blocks." + str(i) + ".cross_attn.k_img.bias"]
274+
cross_k_img_b = rearrange(cross_k_img_b, "(R N D H) -> (N R D H)", R=1, N=40, D=128, H=1)
275+
new_state_dict["decoder.layers." + str(i) + ".cross_attn.linear_k_img.weight"] = cross_k_img_w
276+
new_state_dict["decoder.layers." + str(i) + ".cross_attn.linear_k_img.bias"] = cross_k_img_b
277+
278+
cross_v_img_w = state_dict["blocks." + str(i) + ".cross_attn.v_img.weight"]
279+
cross_v_img_w = rearrange(cross_v_img_w, "(R N D) H -> (N R D) H", R=1, N=40, D=128, H=5120)
280+
cross_v_img_b = state_dict["blocks." + str(i) + ".cross_attn.v_img.bias"]
281+
cross_v_img_b = rearrange(cross_v_img_b, "(R N D H) -> (N R D H)", R=1, N=40, D=128, H=1)
282+
new_state_dict["decoder.layers." + str(i) + ".cross_attn.linear_v_img.weight"] = cross_v_img_w
283+
new_state_dict["decoder.layers." + str(i) + ".cross_attn.linear_v_img.bias"] = cross_v_img_b
284+
246285
# cross_attention o transpose
247286
cross_o_weight = state_dict["blocks." + str(i) + ".cross_attn.o.weight"]
248287
cross_o_weight = rearrange(

examples/wan/convert_checkpoint_mcore2hg.py

Lines changed: 46 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -106,6 +106,16 @@ def save_huggingface_checkpoint(state_dict, save_path):
106106
"time_projection.1.bias",
107107
]
108108
extra_first_part_dict = {
109+
"wan2_1_i2v": [
110+
"img_emb.proj.0.weight",
111+
"img_emb.proj.0.bias",
112+
"img_emb.proj.1.weight",
113+
"img_emb.proj.1.bias",
114+
"img_emb.proj.3.weight",
115+
"img_emb.proj.3.bias",
116+
"img_emb.proj.4.weight",
117+
"img_emb.proj.4.bias",
118+
],
109119
"wan2_2_i2v": []
110120
}
111121
first_part_list = base_first_part_list + extra_first_part_dict.get(model_name, [])
@@ -123,6 +133,11 @@ def save_huggingface_checkpoint(state_dict, save_path):
123133
"blocks.0.cross_attn.norm_q.weight": "decoder.layers.0.cross_attn.q_layernorm.weight",
124134
"blocks.0.cross_attn.norm_k.weight": "decoder.layers.0.cross_attn.k_layernorm.weight",
125135
}
136+
wan2_1_inside_blk_replace_dict = {
137+
"blocks.0.cross_attn.norm_k_img.weight": "decoder.layers.0.cross_attn.k_img_layernorm.weight",
138+
}
139+
if model_name == "wan2_1_i2v":
140+
inside_blk_replace_dict.update(wan2_1_inside_blk_replace_dict)
126141

127142
# Model last part
128143
third_part_dict = {
@@ -206,6 +221,37 @@ def save_huggingface_checkpoint(state_dict, save_path):
206221
new_state_dict["blocks." + str(i) + ".cross_attn.v.weight"] = v_weight
207222
new_state_dict["blocks." + str(i) + ".cross_attn.v.bias"] = v_bias
208223

224+
if model_name == "wan2_1_i2v":
225+
cross_k_img_weight = mcore_dict[
226+
"decoder.layers." + str(i) + ".cross_attn.linear_k_img.weight"
227+
]
228+
cross_k_img_weight = rearrange(
229+
cross_k_img_weight, "(N R D) H -> (R N D) H", R=1, N=40, D=128, H=5120
230+
)
231+
cross_k_img_bias = mcore_dict[
232+
"decoder.layers." + str(i) + ".cross_attn.linear_k_img.bias"
233+
]
234+
cross_k_img_bias = rearrange(
235+
cross_k_img_bias, "(N R D H) -> (R N D H)", R=1, N=40, D=128, H=1
236+
)
237+
new_state_dict["blocks." + str(i) + ".cross_attn.k_img.weight"] = cross_k_img_weight
238+
new_state_dict["blocks." + str(i) + ".cross_attn.k_img.bias"] = cross_k_img_bias
239+
240+
cross_v_img_weight = mcore_dict[
241+
"decoder.layers." + str(i) + ".cross_attn.linear_v_img.weight"
242+
]
243+
cross_v_img_weight = rearrange(
244+
cross_v_img_weight, "(N R D) H -> (R N D) H", R=1, N=40, D=128, H=5120
245+
)
246+
cross_v_img_bias = mcore_dict[
247+
"decoder.layers." + str(i) + ".cross_attn.linear_v_img.bias"
248+
]
249+
cross_v_img_bias = rearrange(
250+
cross_v_img_bias, "(N R D H) -> (R N D H)", R=1, N=40, D=128, H=1
251+
)
252+
new_state_dict["blocks." + str(i) + ".cross_attn.v_img.weight"] = cross_v_img_weight
253+
new_state_dict["blocks." + str(i) + ".cross_attn.v_img.bias"] = cross_v_img_bias
254+
209255
# cross_attention o
210256
cross_o_weight = mcore_dict[
211257
"decoder.layers." + str(i) + ".cross_attn.linear_proj.weight"

examples/wan/convert_wan2.1.sh

Lines changed: 37 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,37 @@
1+
#!/bin/bash
2+
# Copyright 2026 The LoongForge Authors.
3+
# SPDX-License-Identifier: Apache-2.0
4+
5+
if [ $# -eq 0 ]; then
6+
echo "Usage: $0 input \"hg2mcore\" or \"mcore2hg\""
7+
exit 1
8+
fi
9+
input_string=$1
10+
11+
export MEGATRON_PATH=${MEGATRON_PATH:-"/workspace/wan2.1/Loong-Megatron/"}
12+
export LOONGFORGE_PATH=${LOONGFORGE_PATH:-"/workspace/wan2.1/LoongForge"}
13+
export PYTHONPATH=$MEGATRON_PATH:$LOONGFORGE_PATH:$PYTHONPATH
14+
15+
HF_CHECKPOINT_PATH=${HF_CHECKPOINT_PATH:-"/workspace/Wan-AI/Wan2.1-I2V-14B-480P"}
16+
MCORE_CHECKPOINT_PATH=${MCORE_CHECKPOINT_PATH:-"/mnt/cluster/LoongForge/wan2.1/hg2mcore/i2v_14b_480p/Megatron_Release/"}
17+
HG_SAVE_PATH=${HG_SAVE_PATH:-"/mnt/cluster/LoongForge/wan2.1/hg/i2v_14b_480p_dcp/"}
18+
19+
if [ "$input_string" == "hg2mcore" ]; then
20+
echo "convert Wan2.1 I2V weight from huggingface to megatron"
21+
python ./convert_checkpoint_hg2mcore.py \
22+
--save_path="$MCORE_CHECKPOINT_PATH" \
23+
--checkpoint_path="$HF_CHECKPOINT_PATH" \
24+
--num_checkpoints=7 \
25+
--num_layers=40 \
26+
--model_name="wan2_1_i2v"
27+
elif [ "$input_string" == "mcore2hg" ]; then
28+
echo "convert Wan2.1 I2V weight from megatron to huggingface"
29+
python ./convert_checkpoint_mcore2hg.py \
30+
--load_path="$MCORE_CHECKPOINT_PATH" \
31+
--save_path="$HG_SAVE_PATH" \
32+
--num_layers=40 \
33+
--model_name="wan2_1_i2v"
34+
else
35+
echo "Usage: $0 input \"hg2mcore\" or \"mcore2hg\""
36+
exit 1
37+
fi

examples/wan/preprocess.sh

Lines changed: 100 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,100 @@
1+
#!/bin/bash
2+
# Copyright 2026 The LoongForge Authors.
3+
# SPDX-License-Identifier: Apache-2.0
4+
5+
set -euo pipefail
6+
7+
show_help() {
8+
cat <<USAGE
9+
Usage:
10+
bash preprocess.sh wan2.1 [output_path]
11+
bash preprocess.sh wan2.2 [output_path]
12+
bash preprocess.sh --help
13+
14+
Arguments:
15+
wan2.1 Preprocess Wan2.1 I2V data with T5 + VAE + CLIP.
16+
wan2.2 Preprocess Wan2.2 I2V data with T5 + VAE.
17+
output_path Optional output directory.
18+
Default for wan2.1: <dataset_base_parent>/wan_2.1_preprocessed
19+
Default for wan2.2: <dataset_base_parent>/wan_2.2_preprocessed
20+
21+
Optional environment overrides:
22+
LOONGFORGE_ROOT Default: /workspace/wan2.1/LoongForge
23+
WAN21_MODEL_ROOT Default: /workspace/Wan2.1-I2V-14B-480P
24+
WAN22_MODEL_ROOT Default: /workspace/Wan2.2-I2V-A14B
25+
DATASET_BASE_PATH Default: ./dataset/samples
26+
DATASET_METADATA_PATH Default: ./dataset/samples/metadata_100.jsonl
27+
HEIGHT Default: 480 (both versions)
28+
WIDTH Default: 832 (both versions)
29+
NUM_FRAMES Default: 81 (wan2.1) / 49 (wan2.2)
30+
USAGE
31+
}
32+
33+
if [[ $# -lt 1 || $# -gt 2 ]]; then
34+
show_help >&2
35+
exit 1
36+
fi
37+
38+
case "$1" in
39+
-h|--help)
40+
show_help
41+
exit 0
42+
;;
43+
wan2.1|2.1|wan2.2|2.2)
44+
WAN_VERSION="$1"
45+
;;
46+
*)
47+
echo "Unsupported version: $1" >&2
48+
show_help >&2
49+
exit 1
50+
;;
51+
esac
52+
53+
SCRIPT_DIR=$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)
54+
LOONGFORGE_ROOT=${LOONGFORGE_ROOT:-/workspace/wan2.1/LoongForge}
55+
WAN21_MODEL_ROOT=${WAN21_MODEL_ROOT:-/workspace/Wan2.1-I2V-14B-480P}
56+
WAN22_MODEL_ROOT=${WAN22_MODEL_ROOT:-/workspace/Wan2.2-I2V-A14B}
57+
58+
DATASET_BASE_PATH=${DATASET_BASE_PATH:-"${SCRIPT_DIR}/dataset/samples"}
59+
DATASET_METADATA_PATH=${DATASET_METADATA_PATH:-"${DATASET_BASE_PATH}/metadata_100.jsonl"}
60+
DATASET_PARENT_DIR=$(cd "$(dirname "${DATASET_BASE_PATH}")" && pwd)
61+
HEIGHT=${HEIGHT:-480}
62+
WIDTH=${WIDTH:-832}
63+
64+
case "${WAN_VERSION}" in
65+
wan2.1|2.1)
66+
NUM_FRAMES=${NUM_FRAMES:-81}
67+
MODEL_T5=${MODEL_T5:-"${WAN21_MODEL_ROOT}/models_t5_umt5-xxl-enc-bf16.pth"}
68+
MODEL_VAE=${MODEL_VAE:-"${WAN21_MODEL_ROOT}/Wan2.1_VAE.pth"}
69+
MODEL_CLIP=${MODEL_CLIP:-"${WAN21_MODEL_ROOT}/models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"}
70+
TOKENIZER_LOCAL_PATH=${TOKENIZER_LOCAL_PATH:-"${WAN21_MODEL_ROOT}/google/umt5-xxl"}
71+
DEFAULT_OUTPUT_PATH="${DATASET_PARENT_DIR}/wan_2.1_preprocessed"
72+
MODEL_PATHS="${MODEL_T5},${MODEL_VAE},${MODEL_CLIP}"
73+
;;
74+
wan2.2|2.2)
75+
NUM_FRAMES=${NUM_FRAMES:-49}
76+
MODEL_T5=${MODEL_T5:-"${WAN22_MODEL_ROOT}/models_t5_umt5-xxl-enc-bf16.pth"}
77+
MODEL_VAE=${MODEL_VAE:-"${WAN22_MODEL_ROOT}/Wan2.1_VAE.pth"}
78+
TOKENIZER_LOCAL_PATH=${TOKENIZER_LOCAL_PATH:-"${WAN22_MODEL_ROOT}/google/umt5-xxl"}
79+
DEFAULT_OUTPUT_PATH="${DATASET_PARENT_DIR}/wan_2.2_preprocessed"
80+
MODEL_PATHS="${MODEL_T5},${MODEL_VAE}"
81+
;;
82+
esac
83+
84+
OUTPUT_PATH=${2:-${OUTPUT_PATH:-"${DEFAULT_OUTPUT_PATH}"}}
85+
86+
echo "Preprocessing ${WAN_VERSION} dataset"
87+
echo " dataset: ${DATASET_BASE_PATH}"
88+
echo " metadata: ${DATASET_METADATA_PATH}"
89+
echo " output: ${OUTPUT_PATH}"
90+
echo " model_paths: ${MODEL_PATHS}"
91+
echo " resolution: ${HEIGHT}x${WIDTH}, num_frames: ${NUM_FRAMES}"
92+
93+
cd "${SCRIPT_DIR}"
94+
accelerate launch "${LOONGFORGE_ROOT}/examples/wan/wan_preprocess.py" \
95+
--dataset_base_path "${DATASET_BASE_PATH}" \
96+
--dataset_metadata_path "${DATASET_METADATA_PATH}" \
97+
--height "${HEIGHT}" --width "${WIDTH}" --num_frames "${NUM_FRAMES}" \
98+
--model_paths "${MODEL_PATHS}" \
99+
--tokenizer_local_path "${TOKENIZER_LOCAL_PATH}" \
100+
--output_path "${OUTPUT_PATH}"

0 commit comments

Comments
 (0)