LTX-2.3-turbo / app.py
Exosfeer
Fix get_duration arg count mismatch; add show_error=True to launch
9aebeda
Raw
History Blame
22 kB
"""
LTX-2.3 Turbo — ZeroGPU Edition
Generates synchronized audio-video content using Lightricks/LTX-2.3 on
free ZeroGPU hardware via Hugging Face Spaces.
Architecture (following alexnasa/ltx-2-TURBO's proven ZeroGPU pattern):
1. Vendored ltx-core and ltx-pipelines added to sys.path before any imports.
2. Model files downloaded at module startup (CPU, no GPU lease).
3. ModelLedger constructed at module level (CPU-only dataclass, no CUDA init).
4. Text encoder loaded at module level (kept in memory for reuse).
5. DistilledPipeline constructed with gemma_root=None (no text encoder in pipeline).
6. Video encoder and transformer pre-loaded at module level via pipeline cache.
7. @spaces.GPU() on encode_prompt — encodes text, returns .detach().cpu() tensors.
8. @spaces.GPU(duration=callable) on generate_video — runs pipeline with pre-encoded
contexts passed as video_context/audio_context kwargs.
9. FP8 quantization fits the 22B transformer on ZeroGPU's A10G (40GB VRAM).
Based on the official LTX-2 codebase: https://github.com/Lightricks/LTX-2
Architecture inspired by alexnasa/ltx-2-TURBO.
"""
# ───────────────────────────────────────────────────────────────────────────
# 0) Add vendored packages to sys.path BEFORE any ltx imports
# ───────────────────────────────────────────────────────────────────────────
import sys
from pathlib import Path
_here = Path(__file__).parent
sys.path.insert(0, str(_here / "packages" / "ltx-pipelines" / "src"))
sys.path.insert(0, str(_here / "packages" / "ltx-core" / "src"))
# ───────────────────────────────────────────────────────────────────────────
# Standard library & third-party imports
# ───────────────────────────────────────────────────────────────────────────
import logging
import os
import random
import tempfile
import time
import traceback
import gradio as gr
import numpy as np
import spaces
import torch
from huggingface_hub import hf_hub_download, snapshot_download
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
# ───────────────────────────────────────────────────────────────────────────
# LTX imports (from vendored packages)
# ───────────────────────────────────────────────────────────────────────────
from ltx_core.model.video_vae import TilingConfig
from ltx_core.quantization import QuantizationPolicy
from ltx_pipelines.distilled import DistilledPipeline
from ltx_pipelines.utils import ModelLedger
from ltx_pipelines.utils.args import ImageConditioningInput
from ltx_pipelines.utils.helpers import generate_enhanced_prompt
# ───────────────────────────────────────────────────────────────────────────
# Constants
# ───────────────────────────────────────────────────────────────────────────
MAX_SEED = np.iinfo(np.int32).max
LTX_REPO = "Lightricks/LTX-2.3"
GEMMA_REPO = "google/gemma-3-12b-it-qat-q4_0-unquantized"
CKPT_DISTILLED = "ltx-2.3-22b-distilled.safetensors"
CKPT_UPSCALER = "ltx-2.3-spatial-upscaler-x2-1.0.safetensors"
RESOLUTION_PRESETS = {
"16:9 (768x512)": (768, 512),
"1:1 (512x512)": (512, 512),
"9:16 (512x768)": (512, 768),
}
# ───────────────────────────────────────────────────────────────────────────
# 1) Download model files at module startup (CPU, no GPU lease)
# ───────────────────────────────────────────────────────────────────────────
logger.info("Downloading LTX model files...")
checkpoint_path = hf_hub_download(repo_id=LTX_REPO, filename=CKPT_DISTILLED)
logger.info(f" Distilled checkpoint: {checkpoint_path}")
spatial_upsampler_path = hf_hub_download(repo_id=LTX_REPO, filename=CKPT_UPSCALER)
logger.info(f" Upscaler: {spatial_upsampler_path}")
logger.info("Downloading Gemma text encoder...")
HF_TOKEN = os.environ.get("HF_TOKEN")
gemma_root = snapshot_download(repo_id=GEMMA_REPO, token=HF_TOKEN)
logger.info(f" Gemma root: {gemma_root}")
logger.info("All model files ready on disk.")
# ───────────────────────────────────────────────────────────────────────────
# 2) Construct ModelLedger (CPU — no model weights loaded to GPU)
# We construct a separate ledger WITH gemma_root for text encoding.
# ───────────────────────────────────────────────────────────────────────────
logger.info("Constructing ModelLedger (with Gemma for text encoding)...")
fp8_quantization = QuantizationPolicy.fp8_cast()
model_ledger = ModelLedger(
dtype=torch.bfloat16,
device="cuda",
checkpoint_path=checkpoint_path,
gemma_root_path=gemma_root,
spatial_upsampler_path=spatial_upsampler_path,
loras=(),
quantization=fp8_quantization,
)
logger.info("ModelLedger constructed.")
# ───────────────────────────────────────────────────────────────────────────
# 3) Load text encoder at module level (kept in memory for reuse)
# ───────────────────────────────────────────────────────────────────────────
logger.info("Loading Gemma text encoder...")
text_encoder = model_ledger.text_encoder()
logger.info("Text encoder loaded and ready!")
# ───────────────────────────────────────────────────────────────────────────
# 4) Construct DistilledPipeline WITHOUT text encoder (gemma_root=None)
# Text encoding is handled externally via encode_prompt().
# ───────────────────────────────────────────────────────────────────────────
logger.info("Constructing DistilledPipeline (gemma_root=None)...")
pipeline = DistilledPipeline(
device=torch.device("cuda"),
checkpoint_path=checkpoint_path,
spatial_upsampler_path=spatial_upsampler_path,
gemma_root=None,
loras=[],
quantization=fp8_quantization,
)
# ───────────────────────────────────────────────────────────────────────────
# 5) Pre-load video encoder and transformer at module level
# ───────────────────────────────────────────────────────────────────────────
logger.info("Pre-loading video encoder and transformer...")
pipeline._video_encoder = pipeline.model_ledger.video_encoder()
pipeline._transformer = pipeline.model_ledger.transformer()
logger.info("=" * 60)
logger.info("Pipeline fully loaded and ready!")
logger.info("=" * 60)
# ───────────────────────────────────────────────────────────────────────────
# Helpers
# ───────────────────────────────────────────────────────────────────────────
def calc_frames(duration: float, fps: float) -> int:
"""Compute num_frames = 8k + 1, frames >= 9."""
raw = int(duration * fps) + 1
raw = max(raw, 9)
k = (raw - 1 + 7) // 8
return k * 8 + 1
def get_duration(
first_frame,
prompt,
duration,
enhance_prompt,
seed,
randomize_seed,
resolution,
*args,
**kwargs,
):
"""Estimate GPU lease duration for @spaces.GPU(duration=...)."""
dur = float(duration)
if dur <= 2:
return 120
elif dur <= 4:
return 180
else:
return 240
# ───────────────────────────────────────────────────────────────────────────
# Phase 1: Text Encoding (separate GPU lease)
# ───────────────────────────────────────────────────────────────────────────
@spaces.GPU()
def encode_prompt(
prompt: str,
enhance_prompt: bool = True,
input_image=None,
seed: int = 42,
):
"""
Encode prompt using the module-level text_encoder + embeddings_processor.
Returns a dict with video_context and audio_context tensors on CPU.
"""
logger.info(f"[encode_prompt] prompt='{prompt[:80]}...', enhance={enhance_prompt}")
final_prompt = prompt
if enhance_prompt:
final_prompt = generate_enhanced_prompt(
text_encoder=text_encoder,
prompt=prompt,
image_path=input_image if input_image is not None else None,
seed=seed,
)
logger.info(f"[encode_prompt] Enhanced prompt: '{final_prompt[:120]}...'")
with torch.inference_mode():
# Step 1: Get raw hidden states from text encoder
hidden_states, attention_mask = text_encoder.encode(final_prompt)
# Step 2: Process through embeddings processor to get video/audio contexts
embeddings_processor = model_ledger.gemma_embeddings_processor()
result = embeddings_processor.process_hidden_states(
hidden_states, attention_mask
)
del embeddings_processor
video_context = result.video_encoding
audio_context = result.audio_encoding
embedding_data = {
"video_context": video_context.detach().cpu(),
"audio_context": audio_context.detach().cpu(),
"prompt": final_prompt,
}
logger.info("[encode_prompt] Done.")
return embedding_data, final_prompt
# ───────────────────────────────────────────────────────────────────────────
# Phase 2: Video Generation (separate GPU lease, dynamic duration)
# ───────────────────────────────────────────────────────────────────────────
@spaces.GPU(duration=get_duration)
def generate_video(
first_frame,
prompt: str,
duration: float,
enhance_prompt: bool,
seed: int,
randomize_seed: bool,
resolution: str,
progress=gr.Progress(track_tqdm=True),
):
"""
Full generation: encode prompt (nested GPU call) then run pipeline with
pre-encoded contexts.
"""
if not prompt or not prompt.strip():
raise gr.Error("Please enter a prompt.")
current_seed = random.randint(0, MAX_SEED) if randomize_seed else int(seed)
width, height = RESOLUTION_PRESETS.get(resolution, (768, 512))
num_frames = calc_frames(duration, 24.0)
frame_rate = 24.0
logger.info(
f"[generate_video] seed={current_seed}, {width}x{height}, "
f"frames={num_frames}, duration={duration}s, enhance={enhance_prompt}"
)
# --- Handle input image ---
images = []
image_path_for_enhance = None
if first_frame is not None:
tmp_dir = tempfile.mkdtemp()
temp_image_path = os.path.join(tmp_dir, f"input_{int(time.time())}.png")
if hasattr(first_frame, "save"):
first_frame.save(temp_image_path)
else:
from PIL import Image as PILImage
PILImage.open(first_frame).save(temp_image_path)
images = [
ImageConditioningInput(path=temp_image_path, frame_idx=0, strength=1.0)
]
image_path_for_enhance = temp_image_path
t0 = time.time()
try:
# Phase 1: Encode prompt
embeddings, final_prompt = encode_prompt(
prompt=prompt,
enhance_prompt=enhance_prompt,
input_image=image_path_for_enhance,
seed=current_seed,
)
video_context = embeddings["video_context"].to("cuda", non_blocking=True)
audio_context = embeddings["audio_context"].to("cuda", non_blocking=True)
del embeddings
torch.cuda.empty_cache()
# Phase 2: Run pipeline with pre-encoded contexts
with tempfile.NamedTemporaryFile(suffix=".mp4", delete=False) as tmpfile:
output_path = tmpfile.name
with torch.inference_mode():
pipeline(
prompt=prompt,
output_path=output_path,
seed=current_seed,
height=height,
width=width,
num_frames=num_frames,
frame_rate=frame_rate,
images=images,
tiling_config=TilingConfig.default(),
video_context=video_context,
audio_context=audio_context,
)
del video_context, audio_context
torch.cuda.empty_cache()
elapsed = time.time() - t0
logger.info(f"[generate_video] Done in {elapsed:.1f}s")
except torch.cuda.OutOfMemoryError:
elapsed = time.time() - t0
logger.error(f"OOM after {elapsed:.1f}s")
raise gr.Error("Out of GPU memory. Try a shorter duration or lower resolution.")
except Exception as e:
elapsed = time.time() - t0
tb = traceback.format_exc()
logger.error(f"Generation failed after {elapsed:.1f}s:\n{tb}")
raise gr.Error(f"Generation failed: {type(e).__name__}: {e}")
info_text = (
f"Seed: {current_seed}\n"
f"Resolution: {width}x{height} (upscaled from {width // 2}x{height // 2})\n"
f"Frames: {num_frames} @ {int(frame_rate)} fps\n"
f"Duration: {duration}s\n"
f"Pipeline: Distilled 2-stage (8+4 steps, FP8 quantized)\n"
f"Total time: {elapsed:.1f}s\n"
f"Hardware: ZeroGPU"
)
return output_path, info_text, current_seed
# ───────────────────────────────────────────────────────────────────────────
# UI toggle
# ───────────────────────────────────────────────────────────────────────────
def toggle_image(mode: str):
return gr.update(visible=(mode == "Image to Video"))
def on_mode_change(mode: str):
"""Update image visibility and first_frame value based on mode."""
if mode == "Image to Video":
return gr.update(visible=True)
return gr.update(visible=False, value=None)
# ───────────────────────────────────────────────────────────────────────────
# Gradio UI
# ───────────────────────────────────────────────────────────────────────────
CSS = """
.gradio-container { max-width: 1200px !important; }
.header { text-align: center; margin-bottom: 1rem; }
.generate-btn { min-height: 50px; }
"""
with gr.Blocks(css=CSS, title="LTX-2.3 Turbo", theme=gr.themes.Soft()) as demo:
gr.Markdown(
"""
# LTX-2.3 Turbo (ZeroGPU)
Generate synchronized **video + audio** from text or images using
[Lightricks/LTX-2.3](https://huggingface.co/Lightricks/LTX-2.3) —
a 22B parameter DiT-based audio-video foundation model.
Running on **free ZeroGPU** with FP8 quantization.
Distilled pipeline (8+4 denoising steps, two-stage with 2x spatial upscaling).
""",
elem_classes="header",
)
with gr.Row():
# --- Left: Controls ---
with gr.Column(scale=1):
mode = gr.Radio(
["Text to Video", "Image to Video"],
value="Text to Video",
label="Mode",
)
first_frame = gr.Image(
type="pil",
label="Input Image (first frame)",
visible=False,
)
prompt = gr.Textbox(
label="Prompt",
lines=3,
placeholder="Describe the scene, motion, and audio...",
value=(
"A golden retriever puppy plays in fresh snow, "
"tossing it up with its paws, soft winter sunlight, "
"gentle wind sounds and playful barking"
),
)
with gr.Row():
resolution = gr.Dropdown(
choices=list(RESOLUTION_PRESETS.keys()),
value="16:9 (768x512)",
label="Resolution",
)
duration = gr.Slider(1, 5, value=2, step=0.5, label="Duration (sec)")
with gr.Row():
enhance_prompt = gr.Checkbox(value=True, label="Enhance prompt")
randomize_seed = gr.Checkbox(value=True, label="Random seed")
seed = gr.Slider(0, MAX_SEED, value=42, step=1, label="Seed")
generate_btn = gr.Button(
"Generate Video",
variant="primary",
size="lg",
elem_classes="generate-btn",
)
# --- Right: Output ---
with gr.Column(scale=1):
output_video = gr.Video(label="Generated Video", autoplay=True)
run_info = gr.Textbox(label="Generation Info", lines=7, interactive=False)
# --- Events ---
mode.change(fn=on_mode_change, inputs=mode, outputs=[first_frame])
_inputs = [
first_frame,
prompt,
duration,
enhance_prompt,
seed,
randomize_seed,
resolution,
]
_outputs = [output_video, run_info, seed]
generate_btn.click(fn=generate_video, inputs=_inputs, outputs=_outputs)
# --- Examples ---
gr.Examples(
examples=[
[
None,
"Aerial drone shot of a coastal city at sunset, golden light "
"reflecting off glass buildings, gentle ocean waves, seagulls "
"calling, cinematic ambient soundtrack",
3.0,
True,
42,
True,
"16:9 (768x512)",
],
[
None,
"Close-up of a barista pouring latte art in slow motion, "
"steam rising from the cup, coffee shop ambience with soft jazz",
2.0,
True,
123,
True,
"1:1 (512x512)",
],
[
None,
"A cat sits on a windowsill watching rain fall outside, "
"soft indoor lighting, raindrops on glass, gentle rain sounds",
3.0,
True,
7,
True,
"16:9 (768x512)",
],
],
fn=generate_video,
inputs=_inputs,
outputs=_outputs,
cache_examples=False,
label="Example Prompts",
)
gr.Markdown(
"""
---
**Notes:**
- ZeroGPU provides limited GPU time per request. Shorter durations are more reliable.
- Max duration is capped at 5 seconds to stay within GPU time limits.
- FP8 quantization reduces VRAM usage by ~50% with minimal quality impact.
- The 2x spatial upscaler doubles the initial generation resolution.
Built with [Lightricks/LTX-2.3](https://huggingface.co/Lightricks/LTX-2.3)
| [GitHub](https://github.com/Lightricks/LTX-2)
"""
)
if __name__ == "__main__":
demo.launch(show_error=True)