Skip to content

Commit eef1ce5

Browse files
committed
Add video generation infrastructure + LTX Video export
Shared infrastructure: - CoreAIVideoPipeline library: VideoPipeline protocol, VideoConfiguration, VideoGenerationResult, VideoProgress with phase tracking - VideoWriter: MP4 (AVAssetWriter/HEVC), GIF (ImageIO), PNG frame sequence - video-runner CLI skeleton with ArgumentParser LTX Video Python export: - LTXVideoTransformerWrapper with pre-computed 3D RoPE (same pattern as Flux2) - LTXVideoTextEncoderWrapper (T5), VAE encoder/decoder 3D wrappers - compute_ltx_video_rope() for 3-axis temporal+spatial position embeddings - Registry preset, metadata entry, component specs LTX Video (~2B params, Apache 2.0) uses 32x spatial + 8x temporal VAE compression, producing only ~1024 tokens for a 5-second 512x512 video.
1 parent aa3bbf6 commit eef1ce5

10 files changed

Lines changed: 808 additions & 2 deletions

File tree

Package.swift

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -39,6 +39,12 @@ let package = Package(
3939
"CoreAIObjectDetector"
4040
]
4141
),
42+
.library(
43+
name: "CoreAIVideo",
44+
targets: [
45+
"CoreAIVideoPipeline"
46+
]
47+
),
4248
],
4349
dependencies: [
4450
.package(url: "https://github.com/apple/swift-argument-parser", from: "1.2.0"),
@@ -115,6 +121,19 @@ let package = Package(
115121
]
116122
),
117123

124+
// Video Pipeline
125+
.target(
126+
name: "CoreAIVideoPipeline",
127+
dependencies: [
128+
"CoreAIDiffusionPipeline",
129+
"CoreAIShared",
130+
],
131+
path: "swift/Sources/CoreAIVideoPipeline",
132+
swiftSettings: [
133+
.enableUpcomingFeature("MemberImportVisibility")
134+
]
135+
),
136+
118137
// CXGrammar C bridge
119138
.target(
120139
name: "CXGrammar",
@@ -175,6 +194,18 @@ let package = Package(
175194
.enableUpcomingFeature("MemberImportVisibility")
176195
]
177196
),
197+
.executableTarget(
198+
name: "video-runner",
199+
dependencies: [
200+
"CoreAIVideoPipeline",
201+
"CoreAIShared",
202+
.product(name: "ArgumentParser", package: "swift-argument-parser"),
203+
],
204+
path: "swift/Sources/Tools/video-runner",
205+
swiftSettings: [
206+
.enableUpcomingFeature("MemberImportVisibility")
207+
]
208+
),
178209
.executableTarget(
179210
name: "speech-runner",
180211
dependencies: [

python/src/coreai_models/diffusion/components.py

Lines changed: 58 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,16 @@
3232
dummy_flux2_vae_encoder,
3333
dummy_flux2_vae_encoder_half,
3434
)
35+
from coreai_models.diffusion.ltx_video import (
36+
LTXVideoTextEncoderWrapper,
37+
LTXVideoTransformerWrapper,
38+
LTXVideoVAEDecoderWrapper,
39+
LTXVideoVAEEncoderWrapper,
40+
dummy_ltx_video_text_encoder,
41+
dummy_ltx_video_transformer,
42+
dummy_ltx_video_vae_decoder,
43+
dummy_ltx_video_vae_encoder,
44+
)
3545

3646
# ---------------------------------------------------------------------------
3747
# Torch wrappers — thin adapters that extract the tensor we need from the
@@ -362,6 +372,49 @@ def _dummy_sd3_transformer(pipe: Any, batch_size: int = 2) -> tuple[torch.Tensor
362372
ALL_FLUX2_COMPONENTS: list[str] = list(FLUX2_COMPONENTS.keys())
363373

364374

375+
LTX_VIDEO_COMPONENTS: dict[str, ComponentSpec] = {
376+
"transformer": ComponentSpec(
377+
asset_name="Transformer",
378+
input_names=(
379+
"hidden_states",
380+
"encoder_hidden_states",
381+
"timestep",
382+
"encoder_attention_mask",
383+
"rotary_emb_cos",
384+
"rotary_emb_sin",
385+
),
386+
output_names=("output",),
387+
wrapper_fn=lambda p: LTXVideoTransformerWrapper(p.transformer),
388+
dummy_fn=dummy_ltx_video_transformer,
389+
quantizable=True,
390+
),
391+
"text_encoder": ComponentSpec(
392+
asset_name="TextEncoder",
393+
input_names=("input_ids", "attention_mask"),
394+
output_names=("hidden_states",),
395+
wrapper_fn=lambda p: LTXVideoTextEncoderWrapper(p.text_encoder),
396+
dummy_fn=dummy_ltx_video_text_encoder,
397+
quantizable=True,
398+
),
399+
"vae_decoder": ComponentSpec(
400+
asset_name="VAEDecoder",
401+
input_names=("z",),
402+
output_names=("video",),
403+
wrapper_fn=lambda p: LTXVideoVAEDecoderWrapper(p.vae),
404+
dummy_fn=dummy_ltx_video_vae_decoder,
405+
),
406+
"vae_encoder": ComponentSpec(
407+
asset_name="VAEEncoder",
408+
input_names=("video",),
409+
output_names=("latent",),
410+
wrapper_fn=lambda p: LTXVideoVAEEncoderWrapper(p.vae),
411+
dummy_fn=dummy_ltx_video_vae_encoder,
412+
),
413+
}
414+
415+
ALL_LTX_VIDEO_COMPONENTS: list[str] = list(LTX_VIDEO_COMPONENTS.keys())
416+
417+
365418
SD3_COMPONENTS: dict[str, ComponentSpec] = {
366419
"text_encoder": ComponentSpec(
367420
asset_name="TextEncoder",
@@ -408,12 +461,14 @@ def get_component_registry(
408461
Args:
409462
hf_pipe: The loaded HuggingFace pipeline (unused for routing, but
410463
available for future introspection).
411-
pipeline_type: One of "sd", "sd3", or "flux2".
464+
pipeline_type: One of "sd", "sd3", "flux2", or "ltx_video".
412465
"""
413466
if pipeline_type == "flux2":
414467
return FLUX2_COMPONENTS
415468
if pipeline_type == "sd3":
416469
return SD3_COMPONENTS
470+
if pipeline_type == "ltx_video":
471+
return LTX_VIDEO_COMPONENTS
417472
return SD_COMPONENTS
418473

419474

@@ -423,4 +478,6 @@ def get_valid_components(pipeline_type: str) -> list[str]:
423478
return ALL_FLUX2_COMPONENTS
424479
if pipeline_type == "sd3":
425480
return ALL_SD3_COMPONENTS
481+
if pipeline_type == "ltx_video":
482+
return ALL_LTX_VIDEO_COMPONENTS
426483
return ALL_SD_COMPONENTS

0 commit comments

Comments
 (0)