Skip to content

Commit 62f6e69

Browse files
Add prompt-file support for Wan and LTX generation
1 parent 8e3e843 commit 62f6e69

18 files changed

Lines changed: 356 additions & 139 deletions

src/maxdiffusion/configs/base_wan_14b.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -380,6 +380,7 @@ profiler_steps: 10
380380
enable_jax_named_scopes: False
381381

382382
# Generation parameters
383+
prompt_file: ""
383384
prompt: "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
384385
prompt_2: "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
385386
negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"

src/maxdiffusion/configs/base_wan_1_3b.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -333,6 +333,7 @@ profiler_steps: 10
333333
enable_jax_named_scopes: False
334334

335335
# Generation parameters
336+
prompt_file: ""
336337
prompt: "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
337338
prompt_2: "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
338339
negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"

src/maxdiffusion/configs/base_wan_27b.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -353,6 +353,7 @@ profiler_steps: 10
353353
enable_jax_named_scopes: False
354354

355355
# Generation parameters
356+
prompt_file: ""
356357
prompt: "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
357358
prompt_2: "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
358359
negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"

src/maxdiffusion/configs/base_wan_animate.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -343,6 +343,7 @@ profiler_steps: 10
343343
enable_jax_named_scopes: False
344344

345345
# Generation parameters
346+
prompt_file: ""
346347
prompt: "The person from the reference image follows the motion from the driving videos with natural body movement, stable identity, expressive face, cinematic framing, and realistic lighting."
347348
negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
348349
height: 720

src/maxdiffusion/configs/base_wan_i2v_14b.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -345,6 +345,7 @@ profiler_steps: 10
345345
enable_jax_named_scopes: False
346346

347347
# Generation parameters
348+
prompt_file: ""
348349
prompt: "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. They are raising their left arm for a thumbs up. High quality, ultrarealistic detail and breath-taking movie-like camera shot." #LoRA prompt "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. Appearing behind him is a giant, translucent, pink spiritual manifestation (faxiang) that is synchronized with the man's action and pose."
349350
prompt_2: "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot." #LoRA prompt "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. Appearing behind him is a giant, translucent, pink spiritual manifestation (faxiang) that is synchronized with the man's action and pose."
350351
negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"

src/maxdiffusion/configs/base_wan_i2v_27b.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -346,6 +346,7 @@ profiler_steps: 10
346346
enable_jax_named_scopes: False
347347

348348
# Generation parameters
349+
prompt_file: ""
349350
prompt: "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. They are raising their left arm for a thumbs up. High quality, ultrarealistic detail and breath-taking movie-like camera shot." #LoRA prompt "orbit 180 around an astronaut on the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
350351
prompt_2: "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot." #LoRA prompt "orbit 180 around an astronaut on the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
351352
negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"

src/maxdiffusion/configs/ltx2_3_video.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,6 +56,7 @@ use_cross_timestep: true
5656
spatio_temporal_guidance_blocks: [28]
5757
fps: 24
5858
pipeline_type: multi-scale
59+
prompt_file: ""
5960
prompt: "A man in a brightly lit room talks on a vintage telephone. In a low, heavy voice, he says, 'I understand. I won't call again. Goodbye.' He hangs up the receiver and looks down with a sad expression. He holds the black rotary phone to his right ear with his right hand, his left hand holding a rocks glass with amber liquid. He wears a brown suit jacket over a white shirt, and a gold ring on his left ring finger. His short hair is neatly combed, and he has light skin with visible wrinkles around his eyes. The camera remains stationary, focused on his face and upper body. The room is brightly lit by a warm light source off-screen to the left, casting shadows on the wall behind him. The scene appears to be from a dramatic movie."
6061
negative_prompt: "shaky, glitchy, low quality, worst quality, deformed, distorted, disfigured, motion smear, motion artifacts, fused fingers, bad anatomy, weird hand, ugly, transition, static."
6162
height: 512

src/maxdiffusion/configs/ltx2_video.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -63,6 +63,7 @@ spatio_temporal_guidance_blocks: []
6363
noise_scale: 1.0
6464
fps: 24
6565
pipeline_type: multi-scale
66+
prompt_file: ""
6667
prompt: "A man in a brightly lit room talks on a vintage telephone. In a low, heavy voice, he says, 'I understand. I won't call again. Goodbye.' He hangs up the receiver and looks down with a sad expression. He holds the black rotary phone to his right ear with his right hand, his left hand holding a rocks glass with amber liquid. He wears a brown suit jacket over a white shirt, and a gold ring on his left ring finger. His short hair is neatly combed, and he has light skin with visible wrinkles around his eyes. The camera remains stationary, focused on his face and upper body. The room is brightly lit by a warm light source off-screen to the left, casting shadows on the wall behind him. The scene appears to be from a dramatic movie."
6768
negative_prompt: "shaky, glitchy, low quality, worst quality, deformed, distorted, disfigured, motion smear, motion artifacts, fused fingers, bad anatomy, weird hand, ugly, transition, static."
6869
height: 512

src/maxdiffusion/configs/ltx_video.yml

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,7 @@ sampler: "from_checkpoint"
2424

2525
# Generation parameters
2626
pipeline_type: multi-scale
27+
prompt_file: ""
2728
prompt: "A man in a dimly lit room talks on a vintage telephone, hangs up, and looks down with a sad expression. He holds the black rotary phone to his right ear with his right hand, his left hand holding a rocks glass with amber liquid. He wears a brown suit jacket over a white shirt, and a gold ring on his left ring finger. His short hair is neatly combed, and he has light skin with visible wrinkles around his eyes. The camera remains stationary, focused on his face and upper body. The room is dark, lit only by a warm light source off-screen to the left, casting shadows on the wall behind him. The scene appears to be from a movie."
2829
#negative_prompt: "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
2930
height: 512

src/maxdiffusion/generate_ltx2.py

Lines changed: 79 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -302,12 +302,17 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
302302

303303
s0 = time.perf_counter()
304304

305-
# Using global_batch_size_to_train_on to map prompts
306-
prompt = getattr(config, "prompt", "A cat playing piano")
307-
prompt = [prompt] * getattr(config, "global_batch_size_to_train_on", 1)
305+
# Load prompts from prompt_file or default prompt
306+
prompt_file = getattr(config, "prompt_file", "")
307+
default_prompt = getattr(config, "prompt", "A cat playing piano")
308+
prompts = max_utils.load_prompts(prompt_file, default_prompt=default_prompt)
309+
batch_size = getattr(config, "global_batch_size_to_train_on", 1)
310+
is_multi_prompt = len(prompts) > 1 or bool(prompt_file)
308311

309-
negative_prompt = getattr(config, "negative_prompt", "")
310-
negative_prompt = [negative_prompt] * getattr(config, "global_batch_size_to_train_on", 1)
312+
# Using global_batch_size_to_train_on to map prompts
313+
warmup_prompt = [prompts[0]] * batch_size
314+
negative_prompt_str = getattr(config, "negative_prompt", "")
315+
warmup_negative_prompt = [negative_prompt_str] * batch_size
311316

312317
max_logging.log(
313318
f"Num steps: {config.num_inference_steps}, height: {config.height}, width: {config.width}, frames: {config.num_frames}"
@@ -322,6 +327,7 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
322327
max_logging.log(f"hardware: {jax.devices()[0].platform}")
323328
max_logging.log(f"number of devices: {jax.device_count()}")
324329
max_logging.log(f"per_device_batch_size: {config.per_device_batch_size}")
330+
max_logging.log(f"total prompts to generate: {len(prompts)}")
325331
max_logging.log("============================================================")
326332

327333
original_enable_profiler = config.get_keys().get("enable_profiler", False)
@@ -368,7 +374,7 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
368374

369375
max_logging.log(f"🚀 Starting warmup compilation pass ({warmup_steps} steps)...")
370376
with aot_cache.warmup_mode():
371-
_ = call_pipeline(config, pipeline, prompt, negative_prompt)
377+
_ = call_pipeline(config, pipeline, warmup_prompt, warmup_negative_prompt)
372378

373379
aot_cache.save_pending()
374380
config.get_keys()["num_inference_steps"] = original_num_steps
@@ -384,54 +390,79 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
384390

385391
s0 = time.perf_counter()
386392
max_logging.log("🚀 Starting actual full-length generation pass...")
387-
out = call_pipeline(config, pipeline, prompt, negative_prompt)
393+
saved_video_path = []
394+
audio_sample_rate = (
395+
getattr(pipeline.vocoder.config, "output_sampling_rate", 24000)
396+
if getattr(pipeline, "vocoder", None) is not None
397+
else 24000
398+
)
399+
fps = getattr(config, "fps", 24)
400+
audio_format = getattr(config, "audio_format", "s16")
401+
model_name = getattr(config, "model_name", "ltx2") or "ltx2"
402+
model_name_prefix = model_name.replace(".", "_")
403+
gcs_output_path = max_utils.get_gcs_output_path(config)
404+
405+
if not is_multi_prompt:
406+
prompt = [prompts[0]] * batch_size
407+
negative_prompt = [negative_prompt_str] * batch_size
408+
out = call_pipeline(config, pipeline, prompt, negative_prompt)
409+
videos = out.frames if hasattr(out, "frames") else out[0]
410+
audios = out.audio if hasattr(out, "audio") else None
411+
for i in range(len(videos)):
412+
video_path = f"{filename_prefix}{model_name_prefix}_output_{getattr(config, 'seed', 0)}_{i}.mp4"
413+
audio_i = audios[i] if audios is not None else None
414+
export_to_video_with_audio(
415+
video=videos[i],
416+
fps=fps,
417+
audio=audio_i,
418+
audio_sample_rate=audio_sample_rate,
419+
output_path=video_path,
420+
audio_format=audio_format,
421+
)
422+
saved_video_path.append(video_path)
423+
if gcs_output_path:
424+
max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos")
425+
else:
426+
for i in range(0, len(prompts), batch_size):
427+
chunk = prompts[i : i + batch_size]
428+
actual_chunk_len = len(chunk)
429+
if actual_chunk_len < batch_size:
430+
padded_chunk = chunk + [chunk[-1]] * (batch_size - actual_chunk_len)
431+
else:
432+
padded_chunk = chunk
433+
negative_prompt = [negative_prompt_str] * batch_size
434+
435+
out = call_pipeline(config, pipeline, padded_chunk, negative_prompt)
436+
videos = out.frames if hasattr(out, "frames") else out[0]
437+
audios = out.audio if hasattr(out, "audio") else None
438+
for j in range(actual_chunk_len):
439+
prompt_idx = i + j
440+
video_path = f"{filename_prefix}{model_name_prefix}_output_{getattr(config, 'seed', 0)}_{prompt_idx}.mp4"
441+
audio_j = audios[j] if audios is not None else None
442+
export_to_video_with_audio(
443+
video=videos[j],
444+
fps=fps,
445+
audio=audio_j,
446+
audio_sample_rate=audio_sample_rate,
447+
output_path=video_path,
448+
audio_format=audio_format,
449+
)
450+
saved_video_path.append(video_path)
451+
if gcs_output_path:
452+
max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos")
453+
388454
generation_time = time.perf_counter() - s0
389455
max_logging.log(f"generation_time: {generation_time}")
390456
if writer and jax.process_index() == 0:
391457
writer.add_scalar("inference/generation_time", generation_time, global_step=0)
392-
num_devices = jax.device_count()
393-
num_videos = num_devices * config.per_device_batch_size
458+
num_videos = len(saved_video_path)
394459
if num_videos > 0:
395460
generation_time_per_video = generation_time / num_videos
396461
writer.add_scalar("inference/generation_time_per_video", generation_time_per_video, global_step=0)
397462
max_logging.log(f"generation time per video: {generation_time_per_video}")
398463
else:
399464
max_logging.log("Warning: Number of videos is zero, cannot calculate generation_time_per_video.")
400465

401-
# out should have .frames and .audio
402-
videos = out.frames if hasattr(out, "frames") else out[0]
403-
audios = out.audio if hasattr(out, "audio") else None
404-
405-
saved_video_path = []
406-
audio_sample_rate = (
407-
getattr(pipeline.vocoder.config, "output_sampling_rate", 24000)
408-
if getattr(pipeline, "vocoder", None) is not None
409-
else 24000
410-
)
411-
fps = getattr(config, "fps", 24)
412-
413-
# Export videos
414-
for i in range(len(videos)):
415-
model_name = getattr(config, "model_name", "ltx2") or "ltx2"
416-
model_name_prefix = model_name.replace(".", "_")
417-
video_path = f"{filename_prefix}{model_name_prefix}_output_{getattr(config, 'seed', 0)}_{i}.mp4"
418-
audio_i = audios[i] if audios is not None else None
419-
420-
audio_format = getattr(config, "audio_format", "s16")
421-
422-
export_to_video_with_audio(
423-
video=videos[i],
424-
fps=fps,
425-
audio=audio_i,
426-
audio_sample_rate=audio_sample_rate,
427-
output_path=video_path,
428-
audio_format=audio_format,
429-
)
430-
431-
saved_video_path.append(video_path)
432-
if config.output_dir.startswith("gs://"):
433-
max_utils.upload_file_to_gcs(os.path.join(config.output_dir, config.run_name), video_path, subdir="videos")
434-
435466
timing_str = (
436467
f"\n{'=' * 50}\n"
437468
f" TIMING SUMMARY\n"
@@ -481,8 +512,11 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
481512
config.get_keys()["enable_ml_diagnostics"] = False
482513
config.get_keys()["num_inference_steps"] = profiling_steps
483514

515+
profiler_prompt = [prompts[0]] * batch_size
516+
profiler_negative_prompt = [negative_prompt_str] * batch_size
517+
484518
max_logging.log(f"🚀 Warmup for profiling pass ({profiling_steps} steps)...")
485-
_ = call_pipeline(config, pipeline, prompt, negative_prompt)
519+
_ = call_pipeline(config, pipeline, profiler_prompt, profiler_negative_prompt)
486520

487521
config.get_keys()["enable_profiler"] = original_enable_profiler
488522
config.get_keys()["enable_ml_diagnostics"] = original_enable_mld
@@ -491,7 +525,7 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
491525
profiler = max_utils.Profiler(config, session_name=f"denoise_profile_{profiling_steps}_steps")
492526
profiler.start()
493527

494-
_ = call_pipeline(config, pipeline, prompt, negative_prompt)
528+
_ = call_pipeline(config, pipeline, profiler_prompt, profiler_negative_prompt)
495529

496530
profiler.stop()
497531

0 commit comments

Comments
 (0)