make style and quality

This commit is contained in:
SHYuanBest
2026-03-07 13:11:44 +00:00
parent 947c347782
commit f719fc5e91
7 changed files with 56 additions and 89 deletions
+29 -25
View File
@@ -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
+7 -24
View File
@@ -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
View File
@@ -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(
-2
View File
@@ -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
View File
@@ -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