diff --git a/src/maxdiffusion/configs/base_wan_14b.yml b/src/maxdiffusion/configs/base_wan_14b.yml index fd529c15b..837bbe98b 100644 --- a/src/maxdiffusion/configs/base_wan_14b.yml +++ b/src/maxdiffusion/configs/base_wan_14b.yml @@ -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" diff --git a/src/maxdiffusion/configs/base_wan_1_3b.yml b/src/maxdiffusion/configs/base_wan_1_3b.yml index 46c04dfe2..6f5ea10b6 100644 --- a/src/maxdiffusion/configs/base_wan_1_3b.yml +++ b/src/maxdiffusion/configs/base_wan_1_3b.yml @@ -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" diff --git a/src/maxdiffusion/configs/base_wan_27b.yml b/src/maxdiffusion/configs/base_wan_27b.yml index 35b17f9af..bf8e1c740 100644 --- a/src/maxdiffusion/configs/base_wan_27b.yml +++ b/src/maxdiffusion/configs/base_wan_27b.yml @@ -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" diff --git a/src/maxdiffusion/configs/base_wan_animate.yml b/src/maxdiffusion/configs/base_wan_animate.yml index e753df31e..5e9df7d0d 100644 --- a/src/maxdiffusion/configs/base_wan_animate.yml +++ b/src/maxdiffusion/configs/base_wan_animate.yml @@ -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 diff --git a/src/maxdiffusion/configs/base_wan_i2v_14b.yml b/src/maxdiffusion/configs/base_wan_i2v_14b.yml index d3ca82d85..a129ff66c 100644 --- a/src/maxdiffusion/configs/base_wan_i2v_14b.yml +++ b/src/maxdiffusion/configs/base_wan_i2v_14b.yml @@ -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" diff --git a/src/maxdiffusion/configs/base_wan_i2v_27b.yml b/src/maxdiffusion/configs/base_wan_i2v_27b.yml index 43ae00e62..6a28986fc 100644 --- a/src/maxdiffusion/configs/base_wan_i2v_27b.yml +++ b/src/maxdiffusion/configs/base_wan_i2v_27b.yml @@ -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" diff --git a/src/maxdiffusion/configs/ltx2_3_video.yml b/src/maxdiffusion/configs/ltx2_3_video.yml index 0aeaca647..9c9c5432a 100644 --- a/src/maxdiffusion/configs/ltx2_3_video.yml +++ b/src/maxdiffusion/configs/ltx2_3_video.yml @@ -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 diff --git a/src/maxdiffusion/configs/ltx2_video.yml b/src/maxdiffusion/configs/ltx2_video.yml index f953eb8a0..23a4b104a 100644 --- a/src/maxdiffusion/configs/ltx2_video.yml +++ b/src/maxdiffusion/configs/ltx2_video.yml @@ -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 diff --git a/src/maxdiffusion/configs/ltx_video.yml b/src/maxdiffusion/configs/ltx_video.yml index 4b32c65a3..38e83eecb 100644 --- a/src/maxdiffusion/configs/ltx_video.yml +++ b/src/maxdiffusion/configs/ltx_video.yml @@ -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 diff --git a/src/maxdiffusion/generate_ltx2.py b/src/maxdiffusion/generate_ltx2.py index 9790b848c..3cf18cdfd 100644 --- a/src/maxdiffusion/generate_ltx2.py +++ b/src/maxdiffusion/generate_ltx2.py @@ -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}" @@ -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) @@ -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 @@ -384,13 +390,72 @@ 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) + + 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) + 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) @@ -398,40 +463,6 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): 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" @@ -481,8 +512,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 @@ -491,7 +525,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() diff --git a/src/maxdiffusion/generate_ltx_video.py b/src/maxdiffusion/generate_ltx_video.py index 4f66ceb54..bb771424f 100644 --- a/src/maxdiffusion/generate_ltx_video.py +++ b/src/maxdiffusion/generate_ltx_video.py @@ -180,11 +180,16 @@ def run(config): width_padded = ((config.width - 1) // 32 + 1) * 32 num_frames_padded = ((config.num_frames - 2) // 8 + 1) * 8 + 1 padding = calculate_padding(config.height, config.width, height_padded, width_padded) - prompt_enhancement_words_threshold = config.prompt_enhancement_words_threshold - prompt_word_count = len(config.prompt.split()) - enhance_prompt = prompt_enhancement_words_threshold > 0 and prompt_word_count < prompt_enhancement_words_threshold + prompt_file = getattr(config, "prompt_file", "") + prompts = max_utils.load_prompts(prompt_file, default_prompt=config.prompt) + gcs_output_path = max_utils.get_gcs_output_path(config) + prompt_enhancement_words_threshold = getattr(config, "prompt_enhancement_words_threshold", 0) + any_enhance_prompt = any( + prompt_enhancement_words_threshold > 0 and len(prompt.split()) < prompt_enhancement_words_threshold + for prompt in prompts + ) - pipeline = LTXVideoPipeline.from_pretrained(config, enhance_prompt=enhance_prompt) + pipeline = LTXVideoPipeline.from_pretrained(config, enhance_prompt=any_enhance_prompt) if config.pipeline_type == "multi-scale": pipeline = LTXMultiScalePipeline(pipeline) conditioning_media_paths = config.conditioning_media_paths if isinstance(config.conditioning_media_paths, List) else None @@ -206,60 +211,71 @@ def run(config): else None ) - s0 = time.perf_counter() - images = pipeline( - height=height_padded, - width=width_padded, - num_frames=num_frames_padded, - is_video=True, - output_type="pt", - config=config, - enhance_prompt=enhance_prompt, - conditioning_items=conditioning_items, - seed=config.seed, - ) - max_logging.log(f"Compile time: {time.perf_counter() - s0:.1f}s.") - - (pad_left, pad_right, pad_top, pad_bottom) = padding - pad_bottom = -pad_bottom - pad_right = -pad_right - if pad_bottom == 0: - pad_bottom = images.shape[3] - if pad_right == 0: - pad_right = images.shape[4] - images = images[:, :, : config.num_frames, pad_top:pad_bottom, pad_left:pad_right] - output_dir = Path(f"outputs/{datetime.today().strftime('%Y-%m-%d')}") - output_dir.mkdir(parents=True, exist_ok=True) - - for i in range(images.shape[0]): - # Gathering from B, C, F, H, W to C, F, H, W and then permuting to F, H, W, C - video_np = images[i].permute(1, 2, 3, 0).detach().float().numpy() - # Unnormalizing images to [0, 255] range - video_np = (video_np * 255).astype(np.uint8) - fps = config.frame_rate - height, width = video_np.shape[1:3] - # In case a single image is generated - if video_np.shape[0] == 1: - output_filename = get_unique_filename( - f"image_output_{i}", - ".png", - prompt=config.prompt, - resolution=(height, width, config.num_frames), - dir=output_dir, - ) - imageio.imwrite(output_filename, video_np[0]) - else: - output_filename = get_unique_filename( - f"video_output_{i}", - ".mp4", - prompt=config.prompt, - resolution=(height, width, config.num_frames), - dir=output_dir, - ) - # Write video - with imageio.get_writer(output_filename, fps=fps) as video: - for frame in video_np: - video.append_data(frame) + for prompt_idx, current_prompt in enumerate(prompts): + prompt_word_count = len(current_prompt.split()) + enhance_prompt = ( + prompt_enhancement_words_threshold > 0 and prompt_word_count < prompt_enhancement_words_threshold + ) + + s0 = time.perf_counter() + images = pipeline( + prompt=current_prompt, + height=height_padded, + width=width_padded, + num_frames=num_frames_padded, + is_video=True, + output_type="pt", + config=config, + enhance_prompt=enhance_prompt, + conditioning_items=conditioning_items, + seed=config.seed, + ) + max_logging.log(f"Prompt [{prompt_idx + 1}/{len(prompts)}] Generation time: {time.perf_counter() - s0:.1f}s.") + + (pad_left, pad_right, pad_top, pad_bottom) = padding + pad_bottom = -pad_bottom + pad_right = -pad_right + if pad_bottom == 0: + pad_bottom = images.shape[3] + if pad_right == 0: + pad_right = images.shape[4] + images = images[:, :, : config.num_frames, pad_top:pad_bottom, pad_left:pad_right] + output_dir = Path(f"outputs/{datetime.today().strftime('%Y-%m-%d')}") + output_dir.mkdir(parents=True, exist_ok=True) + + for i in range(images.shape[0]): + # Gathering from B, C, F, H, W to C, F, H, W and then permuting to F, H, W, C + video_np = images[i].permute(1, 2, 3, 0).detach().float().numpy() + # Unnormalizing images to [0, 255] range + video_np = (video_np * 255).astype(np.uint8) + fps = config.frame_rate + height, width = video_np.shape[1:3] + # In case a single image is generated + if video_np.shape[0] == 1: + output_filename = get_unique_filename( + f"image_output_{prompt_idx}_{i}" if len(prompts) > 1 else f"image_output_{i}", + ".png", + prompt=current_prompt, + resolution=(height, width, config.num_frames), + dir=output_dir, + ) + imageio.imwrite(output_filename, video_np[0]) + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, str(output_filename), subdir="images") + else: + output_filename = get_unique_filename( + f"video_output_{prompt_idx}_{i}" if len(prompts) > 1 else f"video_output_{i}", + ".mp4", + prompt=current_prompt, + resolution=(height, width, config.num_frames), + dir=output_dir, + ) + # Write video + with imageio.get_writer(output_filename, fps=fps) as video: + for frame in video_np: + video.append_data(frame) + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, str(output_filename), subdir="videos") def main(argv: Sequence[str]) -> None: diff --git a/src/maxdiffusion/generate_wan.py b/src/maxdiffusion/generate_wan.py index b80f46cd4..948bc9ac0 100644 --- a/src/maxdiffusion/generate_wan.py +++ b/src/maxdiffusion/generate_wan.py @@ -122,25 +122,53 @@ def call_pipeline(config, pipeline, prompt, negative_prompt, num_inference_steps def inference_generate_video(config, pipeline, filename_prefix=""): s0 = time.perf_counter() - prompt = [config.prompt] * config.global_batch_size_to_train_on - negative_prompt = [config.negative_prompt] * config.global_batch_size_to_train_on + prompt_file = getattr(config, "prompt_file", "") + prompts = max_utils.load_prompts(prompt_file, default_prompt=config.prompt) + batch_size = config.global_batch_size_to_train_on + is_multi_prompt = len(prompts) > 1 or bool(prompt_file) max_logging.log( f"Num steps: {config.num_inference_steps}, height: {config.height}, width: {config.width}," - f" frames: {config.num_frames}, video: {filename_prefix}" + f" frames: {config.num_frames}, total prompts: {len(prompts)}, video prefix: {filename_prefix}" ) - videos = call_pipeline(config, pipeline, prompt, negative_prompt) + gcs_output_path = max_utils.get_gcs_output_path(config) + saved_video_paths = [] - max_logging.log(f"video {filename_prefix}, compile time: {(time.perf_counter() - s0)}") - for i in range(len(videos)): - video_path = f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" - export_to_video(videos[i], video_path, fps=config.fps) - 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") - # Delete local files to avoid storing too manys videos - max_utils.delete_file(f"./{video_path}") - return + if not is_multi_prompt: + prompt = [prompts[0]] * batch_size + negative_prompt = [config.negative_prompt] * batch_size + videos = call_pipeline(config, pipeline, prompt, negative_prompt) + max_logging.log(f"video {filename_prefix}, generation time: {(time.perf_counter() - s0):.2f}s") + for i in range(len(videos)): + video_path = f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" + export_to_video(videos[i], video_path, fps=config.fps) + saved_video_paths.append(video_path) + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") + max_utils.delete_file(f"./{video_path}") + 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 = [config.negative_prompt] * batch_size + + videos = call_pipeline(config, pipeline, padded_chunk, negative_prompt) + for j in range(actual_chunk_len): + prompt_idx = i + j + video_path = f"{filename_prefix}wan_output_{config.seed}_{prompt_idx}.mp4" + export_to_video(videos[j], video_path, fps=config.fps) + saved_video_paths.append(video_path) + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") + max_utils.delete_file(f"./{video_path}") + max_logging.log(f"all videos {filename_prefix}, total generation time: {(time.perf_counter() - s0):.2f}s") + + return saved_video_paths def maybe_tune_block_sizes(config): @@ -305,13 +333,18 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): original_enable_profiler = config.enable_profiler if "enable_profiler" in config.get_keys() else False config.get_keys()["enable_profiler"] = False + prompt_file = getattr(config, "prompt_file", "") + prompts = max_utils.load_prompts(prompt_file, default_prompt=config.prompt) + batch_size = config.global_batch_size_to_train_on + is_multi_prompt = len(prompts) > 1 or bool(prompt_file) + # Using global_batch_size_to_train_on so not to create more config variables - prompt = [config.prompt] * config.global_batch_size_to_train_on - negative_prompt = [config.negative_prompt] * config.global_batch_size_to_train_on + warmup_prompt = [prompts[0]] * batch_size + warmup_negative_prompt = [config.negative_prompt] * batch_size max_logging.log( f"Num steps: {config.num_inference_steps}, height: {config.height}, width: {config.width}," - f" frames: {config.num_frames}" + f" frames: {config.num_frames}, total prompts: {len(prompts)}" ) # Warmup with 2 denoising steps instead of a full run: step 0 runs the # high-noise transformer and step 1 crosses the boundary to the low-noise @@ -325,7 +358,7 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): # the warmup pays compile time only, never real denoise compute. The # returned videos are garbage by design and are discarded below. with aot_cache.warmup_mode(): - videos = call_pipeline(config, pipeline, prompt, negative_prompt, num_inference_steps=warmup_steps) + videos = call_pipeline(config, pipeline, warmup_prompt, warmup_negative_prompt, num_inference_steps=warmup_steps) if isinstance(videos, tuple): videos, warmup_trace = videos warmup_str = ", ".join(f"{stage}={seconds:.1f}s" for stage, seconds in warmup_trace.items()) @@ -343,6 +376,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("============================================================") compile_time = time.perf_counter() - s0 @@ -351,25 +385,53 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): writer.add_scalar("inference/compile_time", compile_time, global_step=0) s0 = time.perf_counter() - outputs = call_pipeline(config, pipeline, prompt, negative_prompt) - if isinstance(outputs, tuple): - videos, trace = outputs + saved_video_path = [] + gcs_output_path = max_utils.get_gcs_output_path(config) + + if not is_multi_prompt: + prompt = [prompts[0]] * batch_size + negative_prompt = [config.negative_prompt] * batch_size + outputs = call_pipeline(config, pipeline, prompt, negative_prompt) + if isinstance(outputs, tuple): + videos, trace = outputs + else: + videos = outputs + trace = {} + for i in range(len(videos)): + video_path = f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" + export_to_video(videos[i], video_path, fps=config.fps) + saved_video_path.append(video_path) + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") else: - videos = outputs trace = {} + 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 = [config.negative_prompt] * batch_size + + outputs = call_pipeline(config, pipeline, padded_chunk, negative_prompt) + if isinstance(outputs, tuple): + videos, trace = outputs + else: + videos = outputs + for j in range(actual_chunk_len): + prompt_idx = i + j + video_path = f"{filename_prefix}wan_output_{config.seed}_{prompt_idx}.mp4" + export_to_video(videos[j], video_path, fps=config.fps) + 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 - saved_video_path = [] - for i in range(len(videos)): - video_path = f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" - export_to_video(videos[i], video_path, fps=config.fps) - 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") 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( @@ -414,7 +476,9 @@ def run(config, pipeline=None, filename_prefix="", commit_hash=None): os.environ["XLA_FLAGS"] = f"{xla_flags} {new_flags}" max_logging.log(f"Injected XLA_FLAGS for profiling: {new_flags}") - videos = call_pipeline(config, pipeline, prompt, negative_prompt) + profiler_prompt = [prompts[0]] * batch_size + profiler_negative_prompt = [config.negative_prompt] * batch_size + videos = call_pipeline(config, pipeline, profiler_prompt, profiler_negative_prompt) if isinstance(videos, tuple): videos = videos[0] generation_time_with_profiler = time.perf_counter() - s0 diff --git a/src/maxdiffusion/generate_wan_animate.py b/src/maxdiffusion/generate_wan_animate.py index fa253cbe3..1725e1498 100644 --- a/src/maxdiffusion/generate_wan_animate.py +++ b/src/maxdiffusion/generate_wan_animate.py @@ -169,11 +169,19 @@ def run(config): writer.add_scalar("inference/generation_time", generation_time, global_step=0) filename_prefix = "animate_" - os.makedirs(config.output_dir, exist_ok=True) + gcs_output_path = max_utils.get_gcs_output_path(config) + if not gcs_output_path: + os.makedirs(config.output_dir, exist_ok=True) for i, video in enumerate(videos): - video_path = os.path.join(config.output_dir, f"{filename_prefix}wan_output_{config.seed}_{i}.mp4") + video_path = ( + f"{filename_prefix}wan_output_{config.seed}_{i}.mp4" + if gcs_output_path + else os.path.join(config.output_dir, f"{filename_prefix}wan_output_{config.seed}_{i}.mp4") + ) export_to_video(video, video_path, fps=config.fps) max_logging.log(f"Saved video to {video_path}") + if gcs_output_path: + max_utils.upload_file_to_gcs(gcs_output_path, video_path, subdir="videos") if max_utils.profiler_enabled(config): s0 = time.perf_counter() diff --git a/src/maxdiffusion/max_utils.py b/src/maxdiffusion/max_utils.py index 96630cd74..3f308e277 100644 --- a/src/maxdiffusion/max_utils.py +++ b/src/maxdiffusion/max_utils.py @@ -344,6 +344,78 @@ def save_images(config, images): return paths +def get_gcs_output_path(config) -> str: + """Returns the GCS target directory path if output_dir or base_output_directory starts with gs://.""" + output_dir = getattr(config, "output_dir", "") + base_output_dir = getattr(config, "base_output_directory", "") + gcs_root = output_dir if output_dir.startswith("gs://") else ( + base_output_dir if base_output_dir.startswith("gs://") else "" + ) + if not gcs_root: + return "" + run_name = getattr(config, "run_name", "") + return os.path.join(gcs_root, run_name) if run_name else gcs_root + + +def load_prompts(prompt_file_path: str = "", default_prompt: str = "") -> list[str]: + """Loads prompts from a text file (separated by newline) or returns default prompt. + + Supports local files, GCS URIs (gs://bucket/path/to/prompts.txt), and HTTP/HTTPS URLs. + Each line in the file is treated as a separate prompt. Empty lines and whitespace are stripped. + """ + no_prompts_error = "No prompts found. Both prompt_file_path and default_prompt are empty." + if not prompt_file_path: + if default_prompt: + return [default_prompt] + raise ValueError(no_prompts_error) + + prompt_file_path = prompt_file_path.strip() + if not prompt_file_path: + if default_prompt: + return [default_prompt] + raise ValueError(no_prompts_error) + + max_logging.log(f"Loading prompts from file: {prompt_file_path}") + raw_lines = [] + + if prompt_file_path.startswith("gs://"): + try: + bucket_name, prefix_name = parse_gcs_bucket_and_prefix(prompt_file_path) + storage_client = storage.Client() + bucket = storage_client.get_bucket(bucket_name) + blob = bucket.blob(prefix_name) + content = blob.download_as_text() + raw_lines = content.splitlines() + except Exception as e: + max_logging.log(f"Error loading prompts from GCS path '{prompt_file_path}': {e}") + raise + elif prompt_file_path.startswith("http://") or prompt_file_path.startswith("https://"): + try: + import requests + + response = requests.get(prompt_file_path) + response.raise_for_status() + raw_lines = response.text.splitlines() + except Exception as e: + max_logging.log(f"Error downloading prompts from URL '{prompt_file_path}': {e}") + raise + else: + if not os.path.isfile(prompt_file_path): + raise FileNotFoundError(f"Prompt file not found at local path: {prompt_file_path}") + with open(prompt_file_path, "r", encoding="utf-8") as f: + raw_lines = f.readlines() + + prompts = [line.strip() for line in raw_lines if line.strip()] + if not prompts: + if default_prompt: + max_logging.log(f"Warning: Prompt file '{prompt_file_path}' was empty. Falling back to default prompt.") + return [default_prompt] + raise ValueError(f"Prompt file '{prompt_file_path}' contains no valid non-empty prompts.") + + max_logging.log(f"Successfully loaded {len(prompts)} prompt(s) from {prompt_file_path}") + return prompts + + def upload_file_to_gcs(output_dir: str, file_path: str, subdir: str = ""): """Uploads one generated file to {output_dir}/{subdir}/, logging failures. @@ -354,7 +426,7 @@ def upload_file_to_gcs(output_dir: str, file_path: str, subdir: str = ""): parts = path_without_scheme.split("/", 1) bucket_name = parts[0] folder_name = parts[1] if len(parts) > 1 else "" - destination_blob_name = os.path.join(folder_name, subdir, os.path.basename(file_path)) + destination_blob_name = os.path.normpath(os.path.join(folder_name, subdir, os.path.basename(file_path))).lstrip("/") storage_client = storage.Client() bucket = storage_client.bucket(bucket_name) diff --git a/src/maxdiffusion/pipelines/ltx_video/ltx_video_pipeline.py b/src/maxdiffusion/pipelines/ltx_video/ltx_video_pipeline.py index 4aa3baf10..6bd99854f 100644 --- a/src/maxdiffusion/pipelines/ltx_video/ltx_video_pipeline.py +++ b/src/maxdiffusion/pipelines/ltx_video/ltx_video_pipeline.py @@ -658,7 +658,7 @@ def __call__( **kwargs, ): key = jax.random.PRNGKey(seed) - prompt = self.config.prompt + prompt = kwargs.get("prompt", self.config.prompt) is_video = kwargs.get("is_video", False) if prompt is not None and isinstance(prompt, str): batch_size = 1 @@ -1154,6 +1154,7 @@ def __call__( seed: int = 0, enhance_prompt: bool = False, conditioning_items: Optional[List[ConditioningItem]] = None, + prompt: Optional[Union[str, List[str]]] = None, ) -> Any: # first pass original_output_type = output_type @@ -1178,6 +1179,7 @@ def __call__( conditioning_items=conditioning_items, skip_layer_strategy=None, skip_block_list=config.first_pass["skip_block_list"], + prompt=prompt, ) latents = result max_logging.log("first pass done") @@ -1208,6 +1210,7 @@ def __call__( conditioning_items=conditioning_items, skip_layer_strategy=None, skip_block_list=config.second_pass["skip_block_list"], + prompt=prompt, ) if original_output_type != "latent": diff --git a/src/maxdiffusion/pipelines/wan/wan_pipeline.py b/src/maxdiffusion/pipelines/wan/wan_pipeline.py index 80e150ce6..b93599363 100644 --- a/src/maxdiffusion/pipelines/wan/wan_pipeline.py +++ b/src/maxdiffusion/pipelines/wan/wan_pipeline.py @@ -940,7 +940,7 @@ def _decode_latents_to_video(self, latents: jax.Array, trace: Optional[dict] = N trace["vae_decode_tpu"] = time.perf_counter() - t_vae_tpu_start if hasattr(video, "addressable_shards") and len(video.addressable_shards) > 0: - video = np.asarray(video.addressable_shards[0].data) + video = np.concatenate([np.asarray(shard.data) for shard in video.addressable_shards], axis=0) else: video = np.asarray(video) return video diff --git a/src/maxdiffusion/pyconfig.py b/src/maxdiffusion/pyconfig.py index d1121ca3f..ddec32c78 100644 --- a/src/maxdiffusion/pyconfig.py +++ b/src/maxdiffusion/pyconfig.py @@ -97,6 +97,8 @@ class _HyperParameters: def __init__(self, argv: list[str], **kwargs): with open(argv[1], "r", encoding="utf-8") as yaml_file: raw_data_from_yaml = yaml.safe_load(yaml_file) + if "prompt_file" not in raw_data_from_yaml: + raw_data_from_yaml["prompt_file"] = "" raw_data_from_cmd_line = self._load_kwargs(argv) for k in raw_data_from_cmd_line: @@ -206,6 +208,8 @@ def user_init(raw_keys): raw_keys["names_which_can_be_offloaded"] = [] if "offload_encoders" not in raw_keys: raw_keys["offload_encoders"] = False + if "prompt_file" not in raw_keys: + raw_keys["prompt_file"] = "" raw_keys["weights_dtype"] = jax.numpy.dtype(raw_keys["weights_dtype"]) raw_keys["activations_dtype"] = jax.numpy.dtype(raw_keys["activations_dtype"]) diff --git a/src/maxdiffusion/tests/maxdiffusion_utils_test.py b/src/maxdiffusion/tests/maxdiffusion_utils_test.py index 65708494a..43cc07b4c 100644 --- a/src/maxdiffusion/tests/maxdiffusion_utils_test.py +++ b/src/maxdiffusion/tests/maxdiffusion_utils_test.py @@ -25,6 +25,7 @@ from maxdiffusion.max_utils import ( create_device_mesh, get_flash_block_sizes, + load_prompts, ) from maxdiffusion import (FlaxStableDiffusionXLPipeline, FlaxDDIMScheduler, FlaxDDPMScheduler, maxdiffusion_utils) @@ -37,6 +38,12 @@ class MaxDiffusionUtilsTest(unittest.TestCase): def setUp(self): MaxDiffusionUtilsTest.dummy_data = {} + def test_load_prompts_raises_when_no_prompt_source_exists(self): + with self.assertRaisesRegex(ValueError, "No prompts found"): + load_prompts("", "") + with self.assertRaisesRegex(ValueError, "No prompts found"): + load_prompts(" ", "") + def test_get_dummy_wan_inputs_generates_latents_without_pipeline_prepare_latents(self): config = SimpleNamespace(height=64, width=80, num_frames=9, seed=0) pipeline = SimpleNamespace(