diff --git a/src/maxdiffusion/configs/base_wan_14b.yml b/src/maxdiffusion/configs/base_wan_14b.yml index d690b4893..3cb855d63 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 185b01277..053b76a4d 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 2b547dd3d..fbfaf5c29 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 32f95620c..3e0a8c49e 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 c97be88ec..d5e8f3b21 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 b47ee927f..fae138fc0 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 271b07deb..de4424b13 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 d70154e0e..5a24aafb5 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..6b604e0bd 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,75 @@ 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 + 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) @@ -398,40 +466,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 +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 @@ -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() diff --git a/src/maxdiffusion/generate_ltx_video.py b/src/maxdiffusion/generate_ltx_video.py index 4f66ceb54..ad93ece25 100644 --- a/src/maxdiffusion/generate_ltx_video.py +++ b/src/maxdiffusion/generate_ltx_video.py @@ -180,86 +180,94 @@ 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 - - pipeline = LTXVideoPipeline.from_pretrained(config, enhance_prompt=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 - conditioning_start_frames = config.conditioning_start_frames - conditioning_strengths = None - if conditioning_media_paths: - if not conditioning_strengths: - conditioning_strengths = [1.0] * len(conditioning_media_paths) - conditioning_items = ( - prepare_conditioning( - conditioning_media_paths=conditioning_media_paths, - conditioning_strengths=conditioning_strengths, - conditioning_start_frames=conditioning_start_frames, - height=config.height, - width=config.width, - padding=padding, - ) - if conditioning_media_paths - 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) + 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) + + 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 + + pipeline = LTXVideoPipeline.from_pretrained(config, enhance_prompt=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 + conditioning_start_frames = config.conditioning_start_frames + conditioning_strengths = None + if conditioning_media_paths: + if not conditioning_strengths: + conditioning_strengths = [1.0] * len(conditioning_media_paths) + conditioning_items = ( + prepare_conditioning( + conditioning_media_paths=conditioning_media_paths, + conditioning_strengths=conditioning_strengths, + conditioning_start_frames=conditioning_start_frames, + height=config.height, + width=config.width, + padding=padding, + ) + if conditioning_media_paths + 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"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..da623ca33 100644 --- a/src/maxdiffusion/max_utils.py +++ b/src/maxdiffusion/max_utils.py @@ -344,6 +344,77 @@ 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. + """ + if not prompt_file_path: + if default_prompt: + return [default_prompt] + return [] + + prompt_file_path = prompt_file_path.strip() + if not prompt_file_path: + if default_prompt: + return [default_prompt] + return [] + + 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 +425,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/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"])