fix(wan2.2 train): Fix high/low noise timestep schedule. (#270)
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user