Spaces:
Running on Zero
Running on Zero
File size: 12,445 Bytes
e08bfc9 ef045eb e08bfc9 20ade0d e08bfc9 20ade0d e08bfc9 189efc5 e08bfc9 189efc5 e08bfc9 20ade0d e08bfc9 189efc5 e08bfc9 189efc5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 | """OmniVAE T2AV — text to synchronized audio + video.
Thin Gradio wrapper around the reference inference path shipped in
https://github.com/OpenMOSS/OmniVAE (`generation/infer/t2av/t2av_pipeline.py`),
running the released `t2av_recon_distill_avclip` joint checkpoint from
https://huggingface.co/OpenMOSS-Team/OmniVAE.
Defaults mirror the release smoke test documented in `generation/docs/inference.md`
(`validate_checkpoints.sh --cfg 4`): dual CFG, 50 steps, all four guidance
scales at 4.0, 121 frames @ 256x256 @ 24 fps with a 5.04 s waveform.
"""
import os
import sys
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
import spaces # noqa: E402 (must precede torch / CUDA-touching imports)
import gc # noqa: E402
import json # noqa: E402
import logging # noqa: E402
import random # noqa: E402
import tempfile # noqa: E402
import time # noqa: E402
from pathlib import Path # noqa: E402
import gradio as gr # noqa: E402
import torch # noqa: E402
from huggingface_hub import snapshot_download # noqa: E402
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(name)s: %(message)s")
logger = logging.getLogger("omnivae-demo")
MODEL_REPO = "OpenMOSS-Team/OmniVAE"
EXPERIMENT = "t2av_recon_distill_avclip"
# Only the four subtrees the T2AV path touches (~18 GB of the ~64 GB release).
ALLOW_PATTERNS = [
f"models/dit/t2av/{EXPERIMENT}/*",
f"models/dit/t2av/{EXPERIMENT}/**",
"models/text_encoder/Qwen3.5-0.8B-Base/*",
"models/vae/audio_video/recon_distill_avclip/*",
"models/vae/audio_only/recon_distill_avclip_ft_decoder/*",
# `t2av_pipeline._release_root()` requires both `models/` and `eval/` to
# exist under OMNIVAE_RELEASE_ROOT before it will resolve relative paths.
"eval/data/t2av/versebench_minimal/*",
]
logger.info("Downloading OmniVAE release assets ...")
_t0 = time.time()
RELEASE_ROOT = snapshot_download(MODEL_REPO, allow_patterns=ALLOW_PATTERNS, max_workers=8)
logger.info("Release assets ready in %.1fs at %s", time.time() - _t0, RELEASE_ROOT)
os.environ["OMNIVAE_RELEASE_ROOT"] = RELEASE_ROOT
CHECKPOINT_DIR = os.path.join(RELEASE_ROOT, "models", "dit", "t2av", EXPERIMENT)
import torchaudio # noqa: E402
try: # pragma: no cover
import torchcodec # noqa: F401
except ImportError:
# torchaudio >= 2.10 delegates `save` to torchcodec, which is not installed
# (its wheels are pinned to a specific FFmpeg ABI). The pipeline only needs
# a plain 48 kHz WAV write, so route it through soundfile instead.
import numpy as _np # noqa: E402
import soundfile as _sf # noqa: E402
def _save_with_soundfile(uri, src, sample_rate, **_kwargs):
data = src.detach().to("cpu", torch.float32).numpy()
if data.ndim == 2: # torchaudio is (channels, time); soundfile wants (time, channels)
data = data.T
_sf.write(str(uri), _np.ascontiguousarray(data), int(sample_rate))
torchaudio.save = _save_with_soundfile
logger.info("torchcodec unavailable; torchaudio.save patched to use soundfile")
from t2av_pipeline import generate_one_av, load_joint_av_pipeline # noqa: E402
# ---------------------------------------------------------------------------
# Module-scope load. Components are materialised on CPU (mmap / low_cpu_mem_usage
# path) and then moved with `.to("cuda")` so ZeroGPU can intercept the
# placement; the upstream `device="cuda"` path uses `device_map={"": "cuda:0"}`
# plus `torch.cuda.set_device`, neither of which is ZeroGPU-compatible.
# ---------------------------------------------------------------------------
logger.info("Loading T2AV pipeline from %s ...", CHECKPOINT_DIR)
_t0 = time.time()
PIPE = load_joint_av_pipeline(CHECKPOINT_DIR, device="cpu")
logger.info("Pipeline loaded on CPU in %.1fs", time.time() - _t0)
# `load_univae_ckpt` memoises the raw 4.5 GB + 1.5 GB `state_dict.pt` parses via
# an lru_cache; the VAEs are built, so drop them before the ZeroGPU pack step.
from omnivae_generation.trainer.vae.univae import _load_univae_raw # noqa: E402
_load_univae_raw.cache_clear()
gc.collect()
for _name, _module in (
("text_encoder", PIPE.text_encoder),
("video_vae", PIPE.video_vae),
("audio_vae", PIPE.audio_vae),
("joint_model", PIPE.joint_model),
):
_module.to("cuda")
_module.eval()
logger.info("moved %s to cuda", _name)
PIPE.device = torch.device("cuda")
gc.collect()
logger.info("Pipeline ready (checkpoint_step=%s)", PIPE.checkpoint_step)
# Native shapes the released checkpoint was trained / validated at.
NUM_FRAMES = 121
FPS = 24.0
HEIGHT = 256
WIDTH = 256
AUDIO_SECONDS = 5.0417
MAX_SEED = 2**31 - 1
DEFAULT_SEED = 20260508
# Measured on this Space's zero-a10g: 20 steps -> 40.9 s, 50 steps -> 97.6 s
# wall inside `generate` => ~1.89 s/step plus ~3 s fixed (text encode, VAE
# decode of 121 frames, WAV write, ffmpeg mux). Sized with a ~15-20% margin.
STEP_SECONDS = 2.2
BASE_SECONDS = 5.0
def _duration(*args, **kwargs):
"""Reserve GPU time proportional to the requested step count.
Tolerant of partial call shapes: `gr.Examples` invokes `generate` with only
the prompt, and Gradio appends its `Progress` object.
"""
steps = kwargs.get("num_inference_steps")
if steps is None and len(args) >= 3:
steps = args[2]
try:
steps = int(steps)
except (TypeError, ValueError):
steps = 50
return int(min(300, BASE_SECONDS + STEP_SECONDS * steps))
@spaces.GPU(duration=_duration)
def generate(
prompt: str,
negative_prompt: str = "",
num_inference_steps: int = 50,
guidance_scale: float = 4.0,
seed: int = DEFAULT_SEED,
randomize_seed: bool = False,
progress=gr.Progress(track_tqdm=True),
):
if not prompt or not prompt.strip():
raise gr.Error("Please enter a prompt describing the scene and its sound.")
seed = random.randint(0, MAX_SEED) if randomize_seed else int(seed)
guidance_scale = float(guidance_scale)
out_dir = tempfile.mkdtemp(prefix="omnivae_t2av_")
wall_t0 = time.perf_counter()
record = generate_one_av(
PIPE,
prompt=prompt.strip(),
negative_prompt=(negative_prompt or "").strip(),
mode="joint_av",
# Release smoke test (`--cfg 4`) selects BridgeDiT dual CFG (NFE=3)
# and ties all four guidance scales to the same value.
cfg_mode="dual",
num_inference_steps=int(num_inference_steps),
video_guidance_scale=guidance_scale,
audio_guidance_scale=guidance_scale,
cfg_normalization=False,
video_text_guidance=guidance_scale,
video_modality_guidance=guidance_scale,
audio_text_guidance=guidance_scale,
audio_modality_guidance=guidance_scale,
num_frames=NUM_FRAMES,
fps=FPS,
height=HEIGHT,
width=WIDTH,
audio_duration_seconds=AUDIO_SECONDS,
seed=seed,
output_dir=out_dir,
file_stem="omnivae_t2av",
video_quality=8,
# The joint model is trained on task-prefixed prompts; the released
# validation config keeps the prefix on and the duration suffix off.
wrap_task_prefix=True,
task_prefix_kind="t2av",
append_duration_suffix=False,
)
wall = time.perf_counter() - wall_t0
video_path = record.get("av_path") or record.get("video_path")
if not video_path or not Path(video_path).is_file():
raise gr.Error("Generation produced no video file. Check the Space logs.")
if not record.get("av_path"):
logger.warning("ffmpeg mux unavailable; returning silent video.")
audio_path = record.get("audio_path") or None
details = json.dumps(
{
"wrapped_prompt": record.get("wrapped_prompt"),
"seed": record.get("seed"),
"cfg_mode": record.get("cfg_mode"),
"num_inference_steps": record.get("num_inference_steps"),
"guidance_scale": guidance_scale,
"frames": record.get("decoded_num_frames"),
"resolution": f"{record.get('width')}x{record.get('height')}",
"fps": record.get("fps"),
"audio_sample_rate": record.get("sample_rate"),
"denoise_seconds": round(float(record.get("elapsed_s", 0.0)), 2),
"wall_seconds": round(wall, 2),
},
indent=2,
)
logger.info("generate() finished in %.1fs (denoise %.1fs)", wall, record.get("elapsed_s", 0.0))
return video_path, audio_path, seed, details
EXAMPLES = [
# The two prompts the authors ship as their T2AV demo set
# (generation/examples/prompts/t2av_valid.jsonl).
["A street musician plays acoustic guitar on a busy sidewalk while traffic hums in the background."],
["Ocean waves roll onto a sandy beach at sunset with soft wind and distant seabirds."],
# From the released validation config's authored prompt list
# (models/dit/t2av/t2av_recon_distill_avclip/resolved_config.json).
["An old man telling a story to a group of children sitting around him in a park."],
["A young woman laughing while chatting with a friend at a sunny outdoor cafe."],
["A bowl of steaming noodles on a wooden table in a cozy small restaurant."],
]
CSS = """
#col-container { max-width: 900px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
# Gradio 6 moved `theme` / `css` from the Blocks constructor to `launch()`.
with gr.Blocks() as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""
# OmniVAE — Text to Audio + Video
Generate a short clip with **natively synchronized sound** from a single text prompt.
[`OpenMOSS-Team/OmniVAE`](https://huggingface.co/OpenMOSS-Team/OmniVAE) ·
[code](https://github.com/OpenMOSS/OmniVAE)
OmniVAE is a unified audio-video tokenizer; this Space runs the released
`t2av_recon_distill_avclip` joint text-to-audio-video model built on top of it —
two Z-Image diffusion-transformer branches (video + audio) coupled by bridge
cross-attention, so the picture and the soundtrack are denoised together rather
than dubbed afterwards.
Output is the checkpoint's native resolution: **121 frames · 256x256 · 24 fps ·
48 kHz audio (~5 s)**. A default 50-step run takes about 1.5 minutes.
"""
)
prompt = gr.Textbox(
label="Prompt",
placeholder="Describe the scene and what it sounds like…",
lines=3,
)
run_button = gr.Button("Generate", variant="primary")
video_out = gr.Video(label="Audio + video", autoplay=True)
audio_out = gr.Audio(label="Audio track (48 kHz)")
with gr.Accordion("Advanced settings", open=False):
negative_prompt = gr.Textbox(
label="Negative prompt",
value="",
placeholder="Left empty in the reference configuration",
lines=2,
)
num_inference_steps = gr.Slider(
label="Inference steps", minimum=10, maximum=100, step=1, value=50
)
guidance_scale = gr.Slider(
label="Guidance scale (dual CFG — text & cross-modal)",
minimum=1.0,
maximum=10.0,
step=0.5,
value=4.0,
)
with gr.Row():
seed = gr.Slider(
label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=DEFAULT_SEED
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
details_out = gr.Code(label="Run details", language="json")
gr.Examples(
examples=EXAMPLES,
inputs=[prompt],
outputs=[video_out, audio_out, seed, details_out],
fn=generate,
cache_examples=True,
cache_mode="lazy",
)
gr.on(
triggers=[run_button.click, prompt.submit],
fn=generate,
inputs=[prompt, negative_prompt, num_inference_steps, guidance_scale, seed, randomize_seed],
outputs=[video_out, audio_out, seed, details_out],
)
demo.queue(max_size=8).launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)
|