Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions docs/design/inference_schema_parity_inventory.yaml
Original file line number Diff line number Diff line change
Expand Up @@ -518,6 +518,14 @@ surfaces:
true_cfg_scale: request.sampling.true_cfg_scale
boundary_ratio: request.sampling.boundary_ratio
sigmas: request.sampling.sigmas
pyramid_num_inference_steps_list: request.sampling.pyramid_num_inference_steps_list
history_sizes: request.sampling.history_sizes
num_latent_frames_per_chunk: request.sampling.num_latent_frames_per_chunk
keep_first_frame: request.sampling.keep_first_frame
is_skip_first_chunk: request.sampling.is_skip_first_chunk
use_zero_init: request.sampling.use_zero_init
zero_steps: request.sampling.zero_steps
is_amplify_first_chunk: request.sampling.is_amplify_first_chunk
enable_teacache: request.runtime.enable_teacache
save_video: request.output.save_video
return_frames: request.output.return_frames
Expand Down
85 changes: 85 additions & 0 deletions examples/inference/basic/basic_helios_distilled_t2v.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,85 @@
# SPDX-License-Identifier: Apache-2.0
"""Generate one Helios-Distilled T2V chunk through FastVideo's typed API.

Set ``HELIOS_MODEL_PATH`` to a local snapshot to avoid downloading the public
checkpoint again. The 33-frame example is a short integration run; increase
``num_frames`` to 240 for the official eight-chunk default.
"""

from __future__ import annotations

import os

from fastvideo import VideoGenerator
from fastvideo.api import (
EngineConfig,
GenerationRequest,
GeneratorConfig,
OffloadConfig,
OutputConfig,
SamplingConfig,
)

MODEL_PATH = os.getenv("HELIOS_MODEL_PATH", "BestWishYsh/Helios-Distilled")
OUTPUT_PATH = os.getenv(
"HELIOS_OUTPUT_PATH",
"outputs_video/helios/helios_distilled_t2v.mp4",
)
PROMPT = ("A vibrant tropical fish swims gracefully through a colorful coral reef "
"in clear turquoise water, cinematic close-up, fluid motion, vivid detail.")
NEGATIVE_PROMPT = ("Bright tones, overexposed, static, blurred details, subtitles, paintings, "
"images, overall gray, worst quality, low quality, JPEG artifacts, ugly, "
"deformed, disfigured, still picture, messy background.")


def main() -> None:
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=MODEL_PATH,
engine=EngineConfig(
num_gpus=1,
use_fsdp_inference=False,
offload=OffloadConfig(
dit=False,
dit_layerwise=True,
text_encoder=True,
vae=True,
pin_cpu_memory=False,
),
),
))
request = GenerationRequest(
prompt=PROMPT,
negative_prompt=NEGATIVE_PROMPT,
sampling=SamplingConfig(
seed=42,
height=384,
width=640,
num_frames=33,
fps=24,
num_inference_steps=2,
guidance_scale=1.0,
pyramid_num_inference_steps_list=[2, 2, 2],
history_sizes=[16, 2, 1],
num_latent_frames_per_chunk=9,
keep_first_frame=True,
is_skip_first_chunk=False,
use_zero_init=True,
zero_steps=1,
is_amplify_first_chunk=True,
),
output=OutputConfig(
output_path=OUTPUT_PATH,
save_video=True,
return_frames=False,
),
)

try:
generator.generate(request=request)
finally:
generator.shutdown()


if __name__ == "__main__":
main()
60 changes: 60 additions & 0 deletions fastvideo/api/sampling_param.py
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,16 @@ class SamplingParam:
boundary_ratio: float | None = None
sigmas: list[float] | None = None

# Helios autoregressive spatial-pyramid sampling.
pyramid_num_inference_steps_list: list[int] | None = None
history_sizes: list[int] | None = None
num_latent_frames_per_chunk: int = 9
keep_first_frame: bool = True
is_skip_first_chunk: bool = False
use_zero_init: bool = True
zero_steps: int = 1
is_amplify_first_chunk: bool = False

# TeaCache parameters
enable_teacache: bool = False

Expand Down Expand Up @@ -375,6 +385,56 @@ def add_cli_args(parser: Any) -> Any:
default=SamplingParam.boundary_ratio,
help="Boundary timestep ratio",
)
parser.add_argument(
"--pyramid-num-inference-steps-list",
nargs=3,
type=int,
default=SamplingParam.pyramid_num_inference_steps_list,
help="Denoising steps for the three Helios pyramid stages",
)
parser.add_argument(
"--history-sizes",
nargs=3,
type=int,
default=SamplingParam.history_sizes,
help="Long, mid, and short Helios latent history sizes",
)
parser.add_argument(
"--num-latent-frames-per-chunk",
type=int,
default=SamplingParam.num_latent_frames_per_chunk,
help="Helios autoregressive latent frames per chunk",
)
parser.add_argument(
"--keep-first-frame",
action=StoreBoolean,
default=SamplingParam.keep_first_frame,
help="Keep the first Helios latent frame as prefix conditioning",
)
parser.add_argument(
"--is-skip-first-chunk",
action=StoreBoolean,
default=SamplingParam.is_skip_first_chunk,
help="Skip the first Helios autoregressive chunk",
)
parser.add_argument(
"--use-zero-init",
action=StoreBoolean,
default=SamplingParam.use_zero_init,
help="Enable Helios CFG zero initialization when supported",
)
parser.add_argument(
"--zero-steps",
type=int,
default=SamplingParam.zero_steps,
help="Number of Helios CFG zero-initialization steps",
)
parser.add_argument(
"--is-amplify-first-chunk",
action=StoreBoolean,
default=SamplingParam.is_amplify_first_chunk,
help="Use the amplified DMD schedule for the first Helios chunk",
)
parser.add_argument(
"--save-video",
action="store_true",
Expand Down
10 changes: 10 additions & 0 deletions fastvideo/api/schema.py
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,16 @@ class SamplingConfig:
boundary_ratio: float | None = None
sigmas: list[float] | None = None

# Helios autoregressive spatial-pyramid sampling.
pyramid_num_inference_steps_list: list[int] | None = None
history_sizes: list[int] | None = None
num_latent_frames_per_chunk: int = 9
keep_first_frame: bool = True
is_skip_first_chunk: bool = False
use_zero_init: bool = True
zero_steps: int = 1
is_amplify_first_chunk: bool = False


@dataclass
class RequestRuntimeConfig:
Expand Down
5 changes: 3 additions & 2 deletions fastvideo/configs/models/dits/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
from fastvideo.configs.models.dits.flux import FluxDiTConfig
from fastvideo.configs.models.dits.flux_2 import Flux2Config
from fastvideo.configs.models.dits.glm_image import GlmImageDiTConfig
from fastvideo.configs.models.dits.helios import HeliosConfig
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
Expand All @@ -24,6 +25,6 @@
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "FluxDiTConfig", "Flux2Config",
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig",
"StableAudioConfig", "GlmImageDiTConfig", "LingBotWorld2CausalFastVideoConfig", "LingBotVideoConfig",
"MiniMaxH3Config", "ZImageDiTConfig", "MMAudioArchConfig", "MMAudioTransformerConfig"
"StableAudioConfig", "GlmImageDiTConfig", "HeliosConfig", "LingBotWorld2CausalFastVideoConfig",
"LingBotVideoConfig", "MiniMaxH3Config", "ZImageDiTConfig", "MMAudioArchConfig", "MMAudioTransformerConfig"
]
78 changes: 78 additions & 0 deletions fastvideo/configs/models/dits/helios.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,78 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field

from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.platforms import AttentionBackendEnum


def _is_transformer_block(name: str, module) -> bool:
del module
return name.startswith("blocks.") and name.split(".")[-1].isdigit()


@dataclass
class HeliosArchConfig(DiTArchConfig):
"""Architecture fields for Helios-Distilled's history-aware DiT."""

_fsdp_shard_conditions: list = field(default_factory=lambda: [_is_transformer_block])
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
)
param_names_mapping: dict = field(default_factory=dict)
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)

patch_size: tuple[int, int, int] = (1, 2, 2)
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
added_kv_proj_dim: int | None = None
rope_dim: tuple[int, int, int] = (44, 42, 42)
rope_theta: float = 10000.0
guidance_cross_attn: bool = True
zero_history_timestep: bool = True
has_multi_term_memory_patch: bool = True
is_amplify_history: bool = False
history_scale_mode: str = "per_head"

def __post_init__(self) -> None:
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
if not self.cross_attn_norm:
raise ValueError("Helios currently requires cross_attn_norm=True")
if self.qk_norm != "rms_norm_across_heads":
raise ValueError("Helios currently requires qk_norm='rms_norm_across_heads'")
if self.added_kv_proj_dim is not None:
raise ValueError("Helios added_kv_proj_dim variants are not supported")
if not self.guidance_cross_attn:
raise ValueError("Helios currently requires guidance_cross_attn=True")
if not self.zero_history_timestep:
raise ValueError("Helios currently requires zero_history_timestep=True")
if not self.has_multi_term_memory_patch:
raise ValueError("Helios currently requires has_multi_term_memory_patch=True")
if self.is_amplify_history:
raise ValueError("Helios is_amplify_history variants are not supported")
if self.history_scale_mode != "per_head":
raise ValueError("Helios currently requires history_scale_mode='per_head'")
if sum(self.rope_dim) != self.attention_head_dim:
raise ValueError(
f"Helios rope_dim must sum to attention_head_dim, got {self.rope_dim} and {self.attention_head_dim}")
if any(dim % 2 for dim in self.rope_dim):
raise ValueError(f"Helios rope dimensions must be even: {self.rope_dim}")


@dataclass
class HeliosConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HeliosArchConfig)
prefix: str = "Helios"
8 changes: 5 additions & 3 deletions fastvideo/configs/pipelines/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.helios import HeliosPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.kandinsky5 import Kandinsky5DMDConfig, Kandinsky5I2VConfig, Kandinsky5T2VConfig
from fastvideo.configs.pipelines.lingbotworld2 import LingBotWorld2CausalFastI2V480PConfig
Expand All @@ -21,7 +22,8 @@
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "Kandinsky5T2VConfig",
"Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig", "LingBotVideoT2VConfig",
"MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "MMAudioV2AConfig", "get_pipeline_config_cls_from_name"
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HeliosPipelineConfig", "HYWorldConfig",
"Kandinsky5T2VConfig", "Kandinsky5I2VConfig", "Kandinsky5DMDConfig", "LingBotWorld2CausalFastI2V480PConfig",
"LingBotVideoT2VConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "MMAudioV2AConfig",
"get_pipeline_config_cls_from_name"
]
Loading
Loading