diff --git a/app.py b/app.py index cb05ed3..f24994a 100644 --- a/app.py +++ b/app.py @@ -1,20 +1,14 @@ -import os -import subprocess -import sys -import time import tempfile -import zipfile -import torch +import time import gradio as gr import spaces -from diffusers import ( - AutoencoderKLWan, - HeliosPyramidPipeline, - HeliosDMDScheduler -) +import torch + +from diffusers import AutoencoderKLWan, HeliosDMDScheduler, HeliosPyramidPipeline from diffusers.utils import export_to_video, load_image, load_video + # --------------------------------------------------------------------------- # Pre-load model # --------------------------------------------------------------------------- @@ -23,11 +17,7 @@ MODEL_ID = "BestWishYsh/Helios-Distilled" vae = AutoencoderKLWan.from_pretrained(MODEL_ID, subfolder="vae", torch_dtype=torch.float32) scheduler = HeliosDMDScheduler.from_pretrained(MODEL_ID, subfolder="scheduler") pipe = HeliosPyramidPipeline.from_pretrained( - MODEL_ID, - vae=vae, - scheduler=scheduler, - torch_dtype=torch.bfloat16, - is_distilled=True + MODEL_ID, vae=vae, scheduler=scheduler, torch_dtype=torch.bfloat16, is_distilled=True ) pipe.to("cuda") @@ -51,6 +41,7 @@ except Exception: # compiled_transformer = compile_transformer() # spaces.aoti_apply(compiled_transformer, pipe.transformer) + # --------------------------------------------------------------------------- # Generation # --------------------------------------------------------------------------- @@ -104,6 +95,7 @@ def generate_video( info = f"Generated in {elapsed:.1f}s · {num_frames} frames · {height}×{width}" return tmp.name, info + # --------------------------------------------------------------------------- # UI Setup # --------------------------------------------------------------------------- @@ -115,18 +107,19 @@ def update_conditional_visibility(mode): else: return gr.update(visible=False), gr.update(visible=False) + CSS = """ #header { text-align: center; margin-bottom: 1.5em; } #header h1 { font-size: 2.2em; margin-bottom: 0.2em; } .logo { max-height: 100px; margin: 0 auto 10px auto; display: block; } .link-buttons { display: flex; justify-content: center; gap: 15px; margin-top: 10px; } -.link-buttons a { - background-color: #2b3137; - color: #ffffff !important; - padding: 8px 20px; - border-radius: 6px; - text-decoration: none; - font-weight: 600; +.link-buttons a { + background-color: #2b3137; + color: #ffffff !important; + padding: 8px 20px; + border-radius: 6px; + text-decoration: none; + font-weight: 600; font-size: 1em; transition: all 0.2s ease-in-out; box-shadow: 0 2px 4px rgba(0,0,0,0.1); @@ -178,7 +171,7 @@ with gr.Blocks(css=CSS, title="Helios Video Generation", theme=gr.themes.Soft()) "of hard and soft corals in shades of red, orange, and green. The photo captures " "the fish from a slightly elevated angle, emphasizing its lively movements and the " "vivid colors of its surroundings. A close-up shot with dynamic movement." - ) + ), ) with gr.Accordion("Advanced Settings", open=False): with gr.Row(): @@ -200,7 +193,18 @@ with gr.Blocks(css=CSS, title="Helios Video Generation", theme=gr.themes.Soft()) mode.change(fn=update_conditional_visibility, inputs=[mode], outputs=[image_input, video_input]) generate_btn.click( fn=generate_video, - inputs=[mode, prompt, image_input, video_input, height, width, num_frames, num_inference_steps, seed, is_amplify_first_chunk], + inputs=[ + mode, + prompt, + image_input, + video_input, + height, + width, + num_frames, + num_inference_steps, + seed, + is_amplify_first_chunk, + ], outputs=[video_output, info_output], ) diff --git a/helios/diffusers_version/pipeline_helios_diffusers.py b/helios/diffusers_version/pipeline_helios_diffusers.py index 3e9a8c6..99e2dbd 100644 --- a/helios/diffusers_version/pipeline_helios_diffusers.py +++ b/helios/diffusers_version/pipeline_helios_diffusers.py @@ -695,9 +695,7 @@ class HeliosPipeline(DiffusionPipeline, HeliosLoraLoaderMixin): beta = alpha * (1 - ori_sigma) / math.sqrt(gamma) batch_size, channel, num_frames, height, width = latents.shape - noise = self.sample_block_noise( - batch_size, channel, num_frames, height, width, patch_size, device - ) + noise = self.sample_block_noise(batch_size, channel, num_frames, height, width, patch_size, device) noise = noise.to(device=device, dtype=transformer_dtype) latents = alpha * latents + beta * noise # To fix the block artifact @@ -1153,8 +1151,11 @@ class HeliosPipeline(DiffusionPipeline, HeliosLoraLoaderMixin): if not is_enable_stage2: patch_size = self.transformer.config.patch_size - image_seq_len = num_latent_frames_per_chunk * (height // self.vae_scale_factor_spatial) * (width // self.vae_scale_factor_spatial) // ( - patch_size[0] * patch_size[1] * patch_size[2] + image_seq_len = ( + num_latent_frames_per_chunk + * (height // self.vae_scale_factor_spatial) + * (width // self.vae_scale_factor_spatial) + // (patch_size[0] * patch_size[1] * patch_size[2]) ) sigmas = np.linspace(0.999, 0.0, num_inference_steps + 1)[:-1] if sigmas is None else sigmas mu = calculate_shift( diff --git a/helios/diffusers_version/transformer_helios_diffusers.py b/helios/diffusers_version/transformer_helios_diffusers.py index 34d510b..aec1d75 100644 --- a/helios/diffusers_version/transformer_helios_diffusers.py +++ b/helios/diffusers_version/transformer_helios_diffusers.py @@ -558,36 +558,19 @@ class HeliosTransformer3DModel( _cp_plan = { # Input split at attn level and ffn level. "blocks.*.attn1": { - "hidden_states": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), - "rotary_emb": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), + "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), + "rotary_emb": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), }, "blocks.*.attn2": { - "hidden_states": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), + "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), }, "blocks.*.ffn": { - "hidden_states": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), + "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), }, # Output gather at attn level and ffn level. - **{ - f"blocks.{i}.attn1": ContextParallelOutput(gather_dim=1, expected_dims=3) - for i in range(40) - }, - **{ - f"blocks.{i}.attn2": ContextParallelOutput(gather_dim=1, expected_dims=3) - for i in range(40) - }, - **{ - f"blocks.{i}.ffn": ContextParallelOutput(gather_dim=1, expected_dims=3) - for i in range(40) - }, + **{f"blocks.{i}.attn1": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, + **{f"blocks.{i}.attn2": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, + **{f"blocks.{i}.ffn": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, } @register_to_config diff --git a/helios/modules/transformer_helios.py b/helios/modules/transformer_helios.py index 0d7639f..7e1c654 100644 --- a/helios/modules/transformer_helios.py +++ b/helios/modules/transformer_helios.py @@ -964,36 +964,19 @@ class HeliosTransformer3DModel( _cp_plan = { # Input split at attn level and ffn level. "blocks.*.attn1": { - "hidden_states": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), - "rotary_emb": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), + "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), + "rotary_emb": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), }, "blocks.*.attn2": { - "hidden_states": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), + "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), }, "blocks.*.ffn": { - "hidden_states": ContextParallelInput( - split_dim=1, expected_dims=3, split_output=False - ), + "hidden_states": ContextParallelInput(split_dim=1, expected_dims=3, split_output=False), }, # Output gather at attn level and ffn level. - **{ - f"blocks.{i}.attn1": ContextParallelOutput(gather_dim=1, expected_dims=3) - for i in range(40) - }, - **{ - f"blocks.{i}.attn2": ContextParallelOutput(gather_dim=1, expected_dims=3) - for i in range(40) - }, - **{ - f"blocks.{i}.ffn": ContextParallelOutput(gather_dim=1, expected_dims=3) - for i in range(40) - }, + **{f"blocks.{i}.attn1": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, + **{f"blocks.{i}.attn2": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, + **{f"blocks.{i}.ffn": ContextParallelOutput(gather_dim=1, expected_dims=3) for i in range(40)}, } @register_to_config diff --git a/infer_helios.py b/infer_helios.py index 4c06c4d..ec2cb10 100644 --- a/infer_helios.py +++ b/infer_helios.py @@ -5,17 +5,17 @@ import os os.environ["HF_ENABLE_PARALLEL_LOADING"] = "yes" os.environ["HF_PARALLEL_LOADING_WORKERS"] = "8" -import time import argparse -import pandas as pd -from tqdm import tqdm +import time +import pandas as pd import torch import torch.distributed as dist +from tqdm import tqdm + if importlib.util.find_spec("torch_npu") is not None: import torch_npu - from torch_npu.contrib import transfer_to_npu else: torch_npu = None @@ -549,8 +549,8 @@ def main(): interpolation_steps=args.interpolation_steps, interpolate_time_list=interpolate_time_list, ).frames[0] - # elapsed_time = time.time() - start_time - # print(f"Inference time: {elapsed_time:.2f} seconds ({elapsed_time/60:.2f} minutes)") + # elapsed_time = time.time() - start_time + # print(f"Inference time: {elapsed_time:.2f} seconds ({elapsed_time/60:.2f} minutes)") if not args.enable_parallelism or rank == 0: file_count = len( diff --git a/tools/others/convert_ckpt.py b/tools/others/convert_ckpt.py index 850ca5b..39b0855 100644 --- a/tools/others/convert_ckpt.py +++ b/tools/others/convert_ckpt.py @@ -49,14 +49,12 @@ BASE_DIRS = [ for BASE_DIR in BASE_DIRS: index_path = os.path.join(BASE_DIR, "diffusion_pytorch_model.safetensors.index.json") - def apply_rename(key: str) -> str: for old_prefix, new_prefix in RENAME_RULES: if key.startswith(old_prefix): return new_prefix + key[len(old_prefix) :] return key - print(f"[1/3] Reading {index_path} ...") with open(index_path, "r") as f: index = json.load(f) diff --git a/train_helios.py b/train_helios.py index 37baf9d..48aeba3 100644 --- a/train_helios.py +++ b/train_helios.py @@ -257,9 +257,7 @@ def main(args): ) noise_scheduler_copy = copy.deepcopy(noise_scheduler) else: - noise_scheduler = UniPCMultistepScheduler.from_pretrained( - "scripts/accelerate_configs/scheduler_config.json" - ) + noise_scheduler = UniPCMultistepScheduler.from_pretrained("scripts/accelerate_configs/scheduler_config.json") noise_scheduler_copy = FlowMatchEulerDiscreteScheduler(num_train_timesteps=1000) if args.training_config.is_train_dmd: noise_scheduler.config.flow_shift = args.training_config.dmd_timestep_shift