make style and quality
This commit is contained in:
@@ -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],
|
||||
)
|
||||
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
+6
-6
@@ -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(
|
||||
|
||||
@@ -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)
|
||||
|
||||
+1
-3
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user