fix(wan2.2 train): Fix high/low noise timestep schedule. (#270)

This commit is contained in:
flyingshan
2025-08-01 13:38:15 +08:00
committed by GitHub
parent 4dc94d2b95
commit ee8c09416b
2 changed files with 4 additions and 4 deletions
+2 -2
View File
@@ -1480,10 +1480,10 @@ def main():
vae_stream_2 = None
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
if args.boundary_type == "low":
if args.boundary_type == "high":
start_num_idx = 0
train_sampling_steps = int(args.train_sampling_steps * boundary)
elif args.boundary_type == "high":
elif args.boundary_type == "low":
start_num_idx = int(args.train_sampling_steps * boundary)
train_sampling_steps = args.train_sampling_steps - int(args.train_sampling_steps * boundary)
else:
+2 -2
View File
@@ -1488,10 +1488,10 @@ def main():
vae_stream_2 = None
boundary = config['transformer_additional_kwargs'].get('boundary', 0.900)
if args.boundary_type == "low":
if args.boundary_type == "high":
start_num_idx = 0
train_sampling_steps = int(args.train_sampling_steps * boundary)
elif args.boundary_type == "high":
elif args.boundary_type == "low":
start_num_idx = int(args.train_sampling_steps * boundary)
train_sampling_steps = args.train_sampling_steps - int(args.train_sampling_steps * boundary)
else: