Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/base_wan_14b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -380,6 +380,7 @@ profiler_steps: 10
enable_jax_named_scopes: False

# Generation parameters
prompt_file: ""
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."
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."
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"
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/base_wan_1_3b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -333,6 +333,7 @@ profiler_steps: 10
enable_jax_named_scopes: False

# Generation parameters
prompt_file: ""
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."
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."
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"
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/base_wan_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -353,6 +353,7 @@ profiler_steps: 10
enable_jax_named_scopes: False

# Generation parameters
prompt_file: ""
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."
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."
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"
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/base_wan_animate.yml
Original file line number Diff line number Diff line change
Expand Up @@ -343,6 +343,7 @@ profiler_steps: 10
enable_jax_named_scopes: False

# Generation parameters
prompt_file: ""
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."
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"
height: 720
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/base_wan_i2v_14b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -345,6 +345,7 @@ profiler_steps: 10
enable_jax_named_scopes: False

# Generation parameters
prompt_file: ""
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."
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."
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"
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/base_wan_i2v_27b.yml
Original file line number Diff line number Diff line change
Expand Up @@ -346,6 +346,7 @@ profiler_steps: 10
enable_jax_named_scopes: False

# Generation parameters
prompt_file: ""
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."
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."
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"
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/ltx2_3_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,7 @@ use_cross_timestep: true
spatio_temporal_guidance_blocks: [28]
fps: 24
pipeline_type: multi-scale
prompt_file: ""
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."
negative_prompt: "shaky, glitchy, low quality, worst quality, deformed, distorted, disfigured, motion smear, motion artifacts, fused fingers, bad anatomy, weird hand, ugly, transition, static."
height: 512
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/ltx2_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ spatio_temporal_guidance_blocks: []
noise_scale: 1.0
fps: 24
pipeline_type: multi-scale
prompt_file: ""
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."
negative_prompt: "shaky, glitchy, low quality, worst quality, deformed, distorted, disfigured, motion smear, motion artifacts, fused fingers, bad anatomy, weird hand, ugly, transition, static."
height: 512
Expand Down
1 change: 1 addition & 0 deletions src/maxdiffusion/configs/ltx_video.yml
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ sampler: "from_checkpoint"

# Generation parameters
pipeline_type: multi-scale
prompt_file: ""
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."
#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"
height: 512
Expand Down
127 changes: 82 additions & 45 deletions src/maxdiffusion/generate_ltx2.py
Original file line number Diff line number Diff line change
Expand Up @@ -302,12 +302,17 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):

s0 = time.perf_counter()

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

negative_prompt = getattr(config, "negative_prompt", "")
negative_prompt = [negative_prompt] * getattr(config, "global_batch_size_to_train_on", 1)
# Using global_batch_size_to_train_on to map prompts
warmup_prompt = [prompts[0]] * batch_size
negative_prompt_str = getattr(config, "negative_prompt", "")
warmup_negative_prompt = [negative_prompt_str] * batch_size

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

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

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

aot_cache.save_pending()
config.get_keys()["num_inference_steps"] = original_num_steps
Expand All @@ -384,54 +390,82 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):

s0 = time.perf_counter()
max_logging.log("🚀 Starting actual full-length generation pass...")
out = call_pipeline(config, pipeline, prompt, negative_prompt)
saved_video_path = []
audio_sample_rate = (
getattr(pipeline.vocoder.config, "output_sampling_rate", 24000)
if getattr(pipeline, "vocoder", None) is not None
else 24000
)
fps = getattr(config, "fps", 24)
audio_format = getattr(config, "audio_format", "s16")
model_name = getattr(config, "model_name", "ltx2") or "ltx2"
model_name_prefix = model_name.replace(".", "_")
gcs_output_path = max_utils.get_gcs_output_path(config)

last_out = None
if not is_multi_prompt:
prompt = [prompts[0]] * batch_size
negative_prompt = [negative_prompt_str] * batch_size
out = call_pipeline(config, pipeline, prompt, negative_prompt)
last_out = out
Comment on lines +405 to +410

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

medium

The variable last_out is initialized and assigned but never used. The timing summary on line 478 still references out. You can safely remove last_out to keep the code clean. Please also remove the assignment last_out = out on line 438.

Suggested change
last_out = None
if not is_multi_prompt:
prompt = [prompts[0]] * batch_size
negative_prompt = [negative_prompt_str] * batch_size
out = call_pipeline(config, pipeline, prompt, negative_prompt)
last_out = out
if not is_multi_prompt:
prompt = [prompts[0]] * batch_size
negative_prompt = [negative_prompt_str] * batch_size
out = call_pipeline(config, pipeline, prompt, negative_prompt)

videos = out.frames if hasattr(out, "frames") else out[0]
audios = out.audio if hasattr(out, "audio") else None
for i in range(len(videos)):
video_path = f"{filename_prefix}{model_name_prefix}_output_{getattr(config, 'seed', 0)}_{i}.mp4"
audio_i = audios[i] if audios is not None else None
export_to_video_with_audio(
video=videos[i],
fps=fps,
audio=audio_i,
audio_sample_rate=audio_sample_rate,
output_path=video_path,
audio_format=audio_format,
)
saved_video_path.append(video_path)
if gcs_output_path:
max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos")
else:
for i in range(0, len(prompts), batch_size):
chunk = prompts[i : i + batch_size]
actual_chunk_len = len(chunk)
if actual_chunk_len < batch_size:
padded_chunk = chunk + [chunk[-1]] * (batch_size - actual_chunk_len)
else:
padded_chunk = chunk
negative_prompt = [negative_prompt_str] * batch_size

out = call_pipeline(config, pipeline, padded_chunk, negative_prompt)
last_out = out
videos = out.frames if hasattr(out, "frames") else out[0]
audios = out.audio if hasattr(out, "audio") else None
for j in range(actual_chunk_len):
prompt_idx = i + j
video_path = f"{filename_prefix}{model_name_prefix}_output_{getattr(config, 'seed', 0)}_{prompt_idx}.mp4"
audio_j = audios[j] if audios is not None else None
export_to_video_with_audio(
video=videos[j],
fps=fps,
audio=audio_j,
audio_sample_rate=audio_sample_rate,
output_path=video_path,
audio_format=audio_format,
)
saved_video_path.append(video_path)
if gcs_output_path:
max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos")

generation_time = time.perf_counter() - s0
max_logging.log(f"generation_time: {generation_time}")
if writer and jax.process_index() == 0:
writer.add_scalar("inference/generation_time", generation_time, global_step=0)
num_devices = jax.device_count()
num_videos = num_devices * config.per_device_batch_size
num_videos = len(saved_video_path)
if num_videos > 0:
generation_time_per_video = generation_time / num_videos
writer.add_scalar("inference/generation_time_per_video", generation_time_per_video, global_step=0)
max_logging.log(f"generation time per video: {generation_time_per_video}")
else:
max_logging.log("Warning: Number of videos is zero, cannot calculate generation_time_per_video.")

# out should have .frames and .audio
videos = out.frames if hasattr(out, "frames") else out[0]
audios = out.audio if hasattr(out, "audio") else None

saved_video_path = []
audio_sample_rate = (
getattr(pipeline.vocoder.config, "output_sampling_rate", 24000)
if getattr(pipeline, "vocoder", None) is not None
else 24000
)
fps = getattr(config, "fps", 24)

# Export videos
for i in range(len(videos)):
model_name = getattr(config, "model_name", "ltx2") or "ltx2"
model_name_prefix = model_name.replace(".", "_")
video_path = f"{filename_prefix}{model_name_prefix}_output_{getattr(config, 'seed', 0)}_{i}.mp4"
audio_i = audios[i] if audios is not None else None

audio_format = getattr(config, "audio_format", "s16")

export_to_video_with_audio(
video=videos[i],
fps=fps,
audio=audio_i,
audio_sample_rate=audio_sample_rate,
output_path=video_path,
audio_format=audio_format,
)

saved_video_path.append(video_path)
if config.output_dir.startswith("gs://"):
max_utils.upload_file_to_gcs(os.path.join(config.output_dir, config.run_name), video_path, subdir="videos")

timing_str = (
f"\n{'=' * 50}\n"
f" TIMING SUMMARY\n"
Expand Down Expand Up @@ -481,8 +515,11 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None):
config.get_keys()["enable_ml_diagnostics"] = False
config.get_keys()["num_inference_steps"] = profiling_steps

profiler_prompt = [prompts[0]] * batch_size
profiler_negative_prompt = [negative_prompt_str] * batch_size

max_logging.log(f"🚀 Warmup for profiling pass ({profiling_steps} steps)...")
_ = call_pipeline(config, pipeline, prompt, negative_prompt)
_ = call_pipeline(config, pipeline, profiler_prompt, profiler_negative_prompt)

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

_ = call_pipeline(config, pipeline, prompt, negative_prompt)
_ = call_pipeline(config, pipeline, profiler_prompt, profiler_negative_prompt)

profiler.stop()

Expand Down
Loading
Loading