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)