Merge branch 'main' into speed

This commit is contained in:
bubbliiiing
2025-08-01 14:45:57 +08:00
5 changed files with 25 additions and 15 deletions
+1 -1
View File
@@ -67,7 +67,7 @@ V1.0:
| Wan2.1-T2V-1.3B | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-1.3B-InP) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-T2V-1.3B) | Wanxiang 2.1-1.3B text-to-video weights |
| Wan2.1-T2V-14B | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-1.3B-InP) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-T2V-14B) | Wanxiang 2.1-14B text-to-video weights |
| Wan2.1-I2V-14B-480P | [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-14B-InP) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-I2V-14B-480P) | Wanxiang 2.1-14B-480P image-to-video weights |
| Wan2.1-I2V-14B-720P| [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-14B-InP) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-I2V-14B-480P) | Wanxiang 2.1-14B-720P image-to-video weights |
| Wan2.1-I2V-14B-720P| [🤗Link](https://huggingface.co/alibaba-pai/Wan2.1-Fun-14B-InP) | [😄Link](https://www.modelscope.cn/models/Wan-AI/Wan2.1-I2V-14B-720P) | Wanxiang 2.1-14B-720P image-to-video weights |
#### iii. CogVideoX-Fun
+11 -6
View File
@@ -1479,13 +1479,18 @@ def main():
vae_stream_1 = None
vae_stream_2 = None
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
if args.boundary_type == "low":
# Calculate the index we need
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
split_timesteps = args.train_sampling_steps * boundary
differences = torch.abs(noise_scheduler.timesteps - split_timesteps)
closest_index = torch.argmin(differences).item()
print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}")
if args.boundary_type == "high":
start_num_idx = 0
train_sampling_steps = int(args.train_sampling_steps * boundary)
elif args.boundary_type == "high":
start_num_idx = int(args.train_sampling_steps * boundary)
train_sampling_steps = args.train_sampling_steps - int(args.train_sampling_steps * boundary)
train_sampling_steps = closest_index
elif args.boundary_type == "low":
start_num_idx = closest_index
train_sampling_steps = args.train_sampling_steps - closest_index
else:
start_num_idx = 0
train_sampling_steps = args.train_sampling_steps
+11 -6
View File
@@ -1487,13 +1487,18 @@ def main():
vae_stream_1 = None
vae_stream_2 = None
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
if args.boundary_type == "low":
# Calculate the index we need
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
split_timesteps = args.train_sampling_steps * boundary
differences = torch.abs(noise_scheduler.timesteps - split_timesteps)
closest_index = torch.argmin(differences).item()
print(f"The boundary is {boundary} and the boundary_type is {args.boundary_type}. The closest_index we calculate is {closest_index}")
if args.boundary_type == "high":
start_num_idx = 0
train_sampling_steps = int(args.train_sampling_steps * boundary)
elif args.boundary_type == "high":
start_num_idx = int(args.train_sampling_steps * boundary)
train_sampling_steps = args.train_sampling_steps - int(args.train_sampling_steps * boundary)
train_sampling_steps = closest_index
elif args.boundary_type == "low":
start_num_idx = closest_index
train_sampling_steps = args.train_sampling_steps - closest_index
else:
start_num_idx = 0
train_sampling_steps = args.train_sampling_steps
+1 -1
View File
@@ -114,7 +114,7 @@ class Wan2_2Pipeline(DiffusionPipeline):
"""
_optional_components = ["transformer_2"]
model_cpu_offload_seq = "text_encoder->transformer->transformer_2->vae"
model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae"
_callback_tensor_inputs = [
"latents",
+1 -1
View File
@@ -157,7 +157,7 @@ class Wan2_2I2VPipeline(DiffusionPipeline):
"""
_optional_components = ["transformer_2"]
model_cpu_offload_seq = "text_encoder->transformer->transformer_2->vae"
model_cpu_offload_seq = "text_encoder->transformer_2->transformer->vae"
_callback_tensor_inputs = [
"latents",