File size: 18,666 Bytes
3a016a8
186aa49
 
 
 
 
 
 
5029048
186aa49
3a016a8
 
186aa49
299d59b
7db57ff
299d59b
186aa49
9da256f
9e1b9cc
3a016a8
 
3691fc1
186aa49
 
 
 
3a016a8
 
186aa49
c5ee652
87f8144
6d6e37f
 
 
 
c5ee652
 
6d6e37f
 
c5ee652
 
6d6e37f
c5ee652
4adcec8
6d6e37f
c5ee652
6d6e37f
c5ee652
 
6d6e37f
186aa49
50e22ae
186aa49
3a016a8
 
 
186aa49
 
 
 
 
 
 
 
 
 
3a016a8
 
 
 
 
 
 
e6d78bd
 
186aa49
 
 
 
d8b48ff
186aa49
 
 
 
 
 
 
6d6e37f
 
186aa49
 
d8b48ff
 
186aa49
 
 
 
3a016a8
 
 
 
 
 
186aa49
d8b48ff
186aa49
 
 
 
 
 
 
 
 
 
 
3a016a8
186aa49
 
 
 
9da256f
d8b48ff
 
 
 
 
 
 
 
 
186aa49
 
3a016a8
 
6d6e37f
 
 
 
3691fc1
3a016a8
 
 
3691fc1
 
5ca1dce
186aa49
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3a016a8
 
186aa49
 
 
 
 
 
 
 
 
 
 
 
 
 
5029048
bf1199a
7db57ff
 
5029048
186aa49
5029048
186aa49
 
7db57ff
 
 
 
 
 
 
 
 
 
 
 
3a016a8
 
d957826
 
 
7db57ff
d957826
 
 
 
 
 
 
 
 
 
 
9402204
 
3a016a8
 
 
 
 
 
 
9402204
 
456aa27
3a016a8
 
 
 
 
 
 
9402204
 
 
456aa27
0453841
 
3a016a8
0453841
 
186aa49
 
868ff3a
 
456aa27
 
 
868ff3a
5ca1dce
 
3691fc1
 
 
0453841
 
186aa49
 
 
 
 
 
 
 
 
 
456aa27
186aa49
 
299d59b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
456aa27
 
 
 
 
 
 
 
 
186aa49
299d59b
186aa49
299d59b
186aa49
299d59b
186aa49
9da256f
186aa49
 
 
456aa27
 
299d59b
 
 
 
 
 
 
 
 
186aa49
 
 
 
7db57ff
186aa49
 
 
9da256f
 
 
3a016a8
 
9da256f
186aa49
 
456aa27
186aa49
 
299d59b
 
186aa49
 
 
 
 
456aa27
186aa49
 
 
e6d78bd
 
0453841
186aa49
 
299d59b
9da256f
 
868ff3a
456aa27
186aa49
 
299d59b
186aa49
 
d8b4b35
299d59b
 
 
 
 
d8b4b35
 
299d59b
 
942d78c
456aa27
 
 
 
 
 
7db57ff
 
456aa27
d8b4b35
 
299d59b
 
 
 
186aa49
f3a60b2
b84671f
 
299d59b
 
456aa27
 
 
 
299d59b
 
 
 
 
456aa27
 
 
 
 
 
 
 
 
299d59b
 
 
 
 
 
 
186aa49
 
299d59b
 
186aa49
e6d78bd
 
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
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
"""MiniMax-H3 `t2va` / `fl2va`, split deployment — the denoising half."""

from __future__ import annotations

import os
import tempfile
import time
import traceback
from functools import cache

# Before anything that could initialize CUDA: `import spaces` patches `torch.cuda` so the 72 GiB load can happen at
# startup rather than on GPU time.
import spaces
from fastapi.responses import HTMLResponse
from gradio import Request, Server
from gradio.data_classes import FileData

MODEL_REPO = os.environ.get("H3_MODEL_REPO", "MiniMaxAI/MiniMax-H3")
CONDITIONER_SPACE = os.environ.get("H3_CONDITIONER", "multimodalart/qwen3vl-conditioner")
# `pack` places the transformer at startup, `lazy` moves everything on the first GPU call, `offload` hands placement to
# `ComponentsManager.enable_auto_cpu_offload`.
PLACEMENT = os.environ.get("H3_PLACEMENT", "pack").lower()
# cuDNN's fused attention is 10-20% faster than the SDPA default on this pool and needs nothing installed.
ATTENTION = os.environ.get("H3_ATTENTION", "_native_cudnn").lower()
GPU_SIZE = os.environ.get("H3_GPU_SIZE", "xlarge")

# Must stay identical to the conditioner's table: the *label* goes over the wire, so a canvas that half does not know
# is rejected there and surfaces as a failure here.
CANVASES = {
    # 16:9
    "960x544 · 16:9 fast": (544, 960),
    "1024x576 · 16:9 fast": (576, 1024),
    "1152x640 · 16:9": (640, 1152),
    "1280x704 · 16:9": (704, 1280),
    "1344x768 · 16:9 full": (768, 1344),
    # 9:16
    "544x960 · 9:16 fast": (960, 544),
    "640x1152 · 9:16": (1152, 640),
    "768x1344 · 9:16 full": (1344, 768),
    # 1:1
    "544x544 · 1:1 fast": (544, 544),
    "768x768 · 1:1 full": (768, 768),
    # 4:3 / 3:4
    "768x576 · 4:3 fast": (576, 768),
    "1024x768 · 4:3 full": (768, 1024),
    "576x768 · 3:4 fast": (768, 576),
    "768x1024 · 3:4 full": (1024, 768),
    # 21:9
    "1152x512 · 21:9 fast": (512, 1152),
    "1536x672 · 21:9 full": (672, 1536),
}
DEFAULT_CANVAS = "960x544 · 16:9 fast"
FPS, FRAMES_PER_CHUNK, LATENTS_PER_CHUNK = 24, 17, 5
# It is the *snapped* frame count the ceiling has to hold for: 15 s is 360 frames, which rounds up to 362, i.e.
# 15.083 s, and is refused.
MIN_UI_DURATION, MAX_UI_DURATION = 2, 14


def snap_frames(seconds: float) -> int:
    """The frame count MiniMax-H3's video VAE can decode: the next `17 * n + 5` at 24 fps."""
    frames = max(1, round(float(seconds) * FPS))
    while frames % FRAMES_PER_CHUNK != LATENTS_PER_CHUNK:
        frames += 1
    return frames


def lower_duration_floor(seconds: float = MIN_UI_DURATION) -> None:
    """Let the pipeline generate below its 5 s floor. 56 frames (2.33 s) is fine on the released checkpoint."""
    from diffusers.modular_pipelines.minimax_h3.modular_pipeline import MiniMaxH3ModularPipeline

    MiniMaxH3ModularPipeline.min_duration = property(lambda self: float(seconds))


OUTPUT_DIR = os.path.join(tempfile.gettempdir(), "h3-outputs")

PIPE = None
MANAGER = None
LOAD_ERROR: str | None = None
LOADED_IN: float | None = None
LORA_STATUS: str | None = None


def status() -> str:
    if LOAD_ERROR:
        return LOAD_ERROR
    if PIPE is None:
        return f"Loading `{MODEL_REPO}` (transformer + VAEs, 77.3 GB). Watch the Space logs."
    import h3_aoti

    return (
        f"Ready · transformer + VAEs **bfloat16, unquantized** · placement `{PLACEMENT}` · attention `{ATTENTION}` · "
        f"{h3_aoti.status()} · {LORA_STATUS or 'no LoRA'} · loaded in {LOADED_IN:.0f}s · "
        f"conditioner `{CONDITIONER_SPACE}`"
    )


def load_models() -> str | None:
    """Load the denoising half at startup.

    `MiniMaxH3GeneratorBlocks` declares `transformer`, `vae`, `audio_vae`, the two schedulers and `video_processor`,
    so `load_components` fetches exactly those subfolders — `text_encoder/` and `transformer_ref/` are never touched.
    Both autoencoders carry `_keep_in_fp32_modules` over every module and stay float32: a bfloat16 audio VAE decodes
    the soundtrack roughly 20 dB too quiet.
    """
    global PIPE, MANAGER, LOAD_ERROR, LOADED_IN, LORA_STATUS

    if PIPE is not None or LOAD_ERROR is not None:
        return LOAD_ERROR

    started = time.time()
    try:
        import torch
        from diffusers import ComponentsManager

        from h3_split_blocks import MiniMaxH3GeneratorBlocks

        lower_duration_floor()
        manager = ComponentsManager()
        blocks = MiniMaxH3GeneratorBlocks()
        print(f"[gen] loading {[c.name for c in blocks.expected_components]} from {MODEL_REPO} ...", flush=True)
        pipe = blocks.init_pipeline(MODEL_REPO, components_manager=manager, collection="h3")
        pipe.load_components(dtype=torch.bfloat16)

        # Fold the 4-step Turbo LoRA into the bf16 weights before AoTI packages the blocks, so the compiled forward
        # reads weights that already carry the update. `H3_LORA=off` disables.
        import h3_lora

        LORA_STATUS = h3_lora.apply_lora(pipe.transformer)
        if LORA_STATUS:
            print(f"[gen] {LORA_STATUS}", flush=True)

        pipe.transformer.set_attention_backend(ATTENTION)

        # Still startup, still free: an AoTI package carries no weights and opens its archive lazily inside the GPU
        # worker. Off unless `H3_AOTI=1`.
        import h3_aoti

        h3_aoti.maybe_load(pipe.transformer)

        if PLACEMENT == "pack":
            # Scoped to the transformer. `spaces` packs every startup-resident CUDA tensor into a second on-disk copy,
            # and packing all 77.3 GB busts the 150 GB storage quota; the 61.7 GB transformer alone fits. The ~10 GB of
            # fp32 VAEs move on the first GPU call instead.
            pipe.transformer.to("cuda")

        if PLACEMENT == "offload":
            manager.enable_auto_cpu_offload(device="cuda")
            _arm_decode_hooks(pipe)

        PIPE, MANAGER = pipe, manager
        LOADED_IN = time.time() - started
        print(f"[gen] ready in {LOADED_IN:.0f}s", flush=True)
    except Exception as error:
        traceback.print_exc()
        LOAD_ERROR = f"**Loading `{MODEL_REPO}` failed** after {time.time() - started:.0f}s: `{type(error).__name__}: {error}`"
    return LOAD_ERROR


def _arm_decode_hooks(pipe):
    """Make the offload hooks fire for the two VAEs.

    `enable_auto_cpu_offload` wraps `forward`, and the decode blocks call `vae.decode(...)` directly, so the hook
    never runs and the VAE is still on the host when the latents arrive on the card.
    """
    for name in ("vae", "audio_vae"):
        module = getattr(pipe, name)
        inner = module.decode

        def armed(*args, _module=module, _decode=inner, **kwargs):
            hook = getattr(_module, "_hf_hook", None)
            if hook is not None:
                hook.pre_forward(_module)
            return _decode(*args, **kwargs)

        module.decode = armed


@cache
def conditioner():
    """The other half, over the gradio API. Used only when the caller's token could not be extracted; the booking is
    then billed to this Space's pod IP and its small shared quota."""
    from gradio_client import Client

    return Client(CONDITIONER_SPACE)


def conditioner_client(ip_token):
    """A conditioner client billed to the caller. `LocalContext`-based token forwarding is not reliable in Server
    mode, so the `x-ip-token` header is extracted from the incoming request and passed explicitly (per the gradio
    ZeroGPU docs); a per-request Client is cheap next to a 45s encode."""
    if not ip_token:
        return conditioner()
    from gradio_client import Client

    return Client(CONDITIONER_SPACE, headers={"x-ip-token": ip_token})


def encode_remote(prompt, image_path, last_image_path, canvas, num_frames, rewrite_prompt=False, ip_token=None):
    """`/encode` on the conditioner Space: a safetensors file holding `prompt_embeds` + `text_token_tags`, with the
    resolved `height` / `width` / `num_frames` in its metadata, plus the plan. `canvas` is the label."""
    from gradio_client import handle_file
    from safetensors import safe_open

    path, plan = conditioner_client(ip_token).predict(
        prompt=prompt,
        image_path=handle_file(image_path) if image_path else None,
        last_image_path=handle_file(last_image_path) if last_image_path else None,
        canvas=canvas,
        num_frames=num_frames,
        rewrite_prompt=bool(rewrite_prompt),
        api_name="/encode",
    )
    with safe_open(path, framework="pt") as handle:
        metadata = handle.metadata()
        return handle.get_tensor("prompt_embeds"), handle.get_tensor("text_token_tags"), metadata, plan


# Seconds of GPU one request needs, from the packed video rows it is about to denoise: linear in the rows for the
# matmuls, quadratic for the attention, against the AoTI block package this Space runs.
_DUR_B, _DUR_C = 1.1745e-4, 3.8396e-9
# The two resident decoders and the mux, which scale with the output rather than with the step count.
_DECODE_BASE, _DECODE_PER_DEFAULT_CANVAS, _DEFAULT_CANVAS_PIXELS = 15, 15, 960 * 544 * 124
# `pack` mode: only the ~10 GB of VAEs move on a cold worker.
_PLACEMENT_ALLOWANCE, _PAD = 12, 10


def get_duration(prompt_embeds, text_token_tags, image, last_image, height, width, num_frames, steps, seed, lora="larry", *a, **k):
    height, width, num_frames, steps = int(height), int(width), int(num_frames), int(steps)
    latent_frames = (num_frames - LATENTS_PER_CHUNK) // FRAMES_PER_CHUNK * LATENTS_PER_CHUNK + 2
    patches = (height // 32) * (width // 32)
    rows = latent_frames * patches + (int(image is not None) + int(last_image is not None)) * patches
    denoise = steps * (_DUR_B * rows + _DUR_C * rows**2)
    decode = _DECODE_BASE + _DECODE_PER_DEFAULT_CANVAS * (height * width * num_frames) / _DEFAULT_CANVAS_PIXELS
    return max(60, int(denoise + decode) + _PLACEMENT_ALLOWANCE + _PAD)


@spaces.GPU(duration=get_duration, size=GPU_SIZE)
def _generate(prompt_embeds, text_token_tags, image, last_image, height, width, num_frames, steps, seed, lora="larry"):
    """The only thing on GPU time: the packed-sequence denoise loop and the two decoders.

    Only the three generated outputs come back — a `@spaces.GPU` return crosses a process boundary by pickling, and
    the full `PipelineState` still holds the packed latents, the rotary grid and the row indices on the card.
    """
    import torch

    import h3_lora

    # Fold the requested LoRA in place (a no-op when the state already matches). AoTI blocks read the same
    # live storage, so the compiled forward carries the switch too.
    active_lora = h3_lora.set_active(PIPE.transformer, lora)

    if PLACEMENT == "lazy":
        PIPE.to("cuda")
    elif PLACEMENT == "pack":
        PIPE.vae.to("cuda")
        PIPE.audio_vae.to("cuda")

    state = PIPE(
        prompt_embeds=prompt_embeds.to("cuda"),
        text_token_tags=text_token_tags,
        image=image,
        last_image=last_image,
        height=height,
        width=width,
        num_frames=num_frames,
        num_inference_steps=int(steps),
        generator=torch.Generator("cpu").manual_seed(int(seed)),
    )
    return state.get("videos")[0], state.get("audio")[0].cpu(), state.get("sampling_rate"), active_lora


def _fit_keyframe(image_path, current_canvas):
    """Cover-crop an uploaded keyframe to the closest supported aspect ratio and pick that ratio's smallest
    (fastest) canvas, unless the caller already picked a matching ratio. Returns `(image_path, canvas_label)`."""
    from PIL import Image as _Image

    img = _Image.open(image_path)
    aspect = img.width / img.height
    fastest = {}
    for label, (h, w) in CANVASES.items():
        r = w / h
        if r not in fastest or w * h < fastest[r][1][0] * fastest[r][1][1]:
            fastest[r] = (label, (h, w))
    ratio = min(fastest, key=lambda r: abs(r - aspect))
    label, (h, w) = fastest[ratio]

    cur_h, cur_w = CANVASES[current_canvas]
    if abs(cur_w / cur_h - aspect) <= abs(ratio - aspect):
        label = current_canvas
        h, w = cur_h, cur_w

    target = w / h
    if abs(img.width / img.height - target) > 1e-3:
        if img.width / img.height > target:
            new_w = int(img.height * target)
            left = (img.width - new_w) // 2
            img = img.crop((left, 0, left + new_w, img.height))
        else:
            new_h = int(img.width / target)
            top = (img.height - new_h) // 2
            img = img.crop((0, top, img.width, top + new_h))
        img.save(image_path)
    return image_path, label


def _resolve_lora(lora, use_lora) -> str:
    """`lora` (`larry` / `lightx` / `off`) wins; the legacy `use_lora` bool maps onto `larry` / `off`."""
    if isinstance(lora, str) and lora in ("larry", "lightx", "off"):
        return lora
    return "larry" if use_lora else "off"


def generate(prompt, image_path=None, last_image_path=None, canvas=DEFAULT_CANVAS, duration=5, steps=6, seed=42, upsample=False, use_lora=True, lora="", ip_token=None):
    """One request. `upsample`/`use_lora` keep their defaults so a positional API client that predates them is unaffected."""
    if LOAD_ERROR:
        raise Exception(LOAD_ERROR)
    if PIPE is None:
        raise Exception("The denoiser is still loading.")
    if not prompt or not prompt.strip():
        raise Exception("MiniMax-H3 always takes a prompt, keyframes or not.")

    from PIL import Image, ImageOps

    from diffusers.utils import encode_video

    lora = _resolve_lora(lora, use_lora)

    # Server mode: keyframes arrive as FileData dicts, and the cover-crop / canvas-fit that used to be an upload
    # event in the Blocks UI runs here instead, so API callers get the same treatment.
    first = image_path["path"] if isinstance(image_path, dict) else image_path
    last = last_image_path["path"] if isinstance(last_image_path, dict) else last_image_path
    if first:
        first, canvas = _fit_keyframe(first, canvas)
    if last:
        last, canvas = _fit_keyframe(last, canvas)

    num_frames = snap_frames(duration)

    conditioned = time.time()
    prompt_embeds, text_token_tags, metadata, plan = encode_remote(
        prompt, first, last, canvas, num_frames, rewrite_prompt=upsample, ip_token=ip_token
    )
    condition_seconds = time.time() - conditioned
    height, width, num_frames = (int(metadata[key]) for key in ("height", "width", "num_frames"))
    refined = plan.get("refined_prompt") or ""

    def keyframe(path):
        # The conditioning latents encoded here have to be of the image the conditioner looked at, which it prepares
        # exactly this way.
        return ImageOps.exif_transpose(Image.open(path)).convert("RGB") if path else None

    started = time.time()
    frames, audio, sampling_rate, active_lora = _generate(
        prompt_embeds,
        text_token_tags,
        keyframe(first),
        keyframe(last),
        height,
        width,
        num_frames,
        steps,
        seed,
        lora,
    )
    generate_seconds = time.time() - started

    os.makedirs(OUTPUT_DIR, exist_ok=True)
    path = os.path.join(OUTPUT_DIR, f"h3-{int(time.time() * 1000)}.mp4")
    encode_video(frames, fps=FPS, output_path=path, audio=audio, audio_sample_rate=sampling_rate)

    report = (
        f"{width}x{height} · {num_frames} frames ({num_frames / FPS:.3f} s) · {int(steps)} steps · "
        f"conditioner {condition_seconds:.0f}s ({plan['num_text_tokens']} tokens"
        f"{', upsampled' if refined else ''}) · "
        f"denoise + decode {generate_seconds:.0f}s ({generate_seconds / int(steps):.1f} s/step) · "
        f"turbo LoRA {active_lora} · seed {int(seed)}"
    )
    print(f"[gen] {report}", flush=True)
    return FileData(path=path), report, refined



# ======================================================================
# Server mode: Gradio's API engine (queue, SSE, concurrency, ZeroGPU,
# gradio_client) under a fully custom studio frontend (index.html).
# ======================================================================
app = Server(title="MiniMax-H3 Studio")


@app.api(name="generate")
def _generate_api(prompt: str, image_path: FileData | None = None, last_image_path: FileData | None = None,
                  canvas: str = DEFAULT_CANVAS, duration: float = 5, steps: int = 6, seed: float = 42,
                  upsample: bool = False, use_lora: bool = True, lora: str = "", request: Request = None) -> tuple[FileData, str, str]:
    """Generate a video with a synchronized soundtrack. Returns (video, report, refined prompt).

    `lora` selects the turbo LoRA: `larry` (default), `lightx`, or `off`. The legacy `use_lora` bool still works
    when `lora` is empty.
    """
    # `request` is injected by the event system, not an API input; its x-ip-token bills the conditioner to the caller.
    ip_token = request.headers.get("x-ip-token") if request is not None else None
    return generate(prompt, image_path, last_image_path, canvas, duration, steps, seed, upsample, use_lora, lora, ip_token=ip_token)


@app.get("/status")
def studio_status():
    """Polled by the frontend: is the denoiser ready, and the human-readable status line."""
    return {"ready": PIPE is not None and LOAD_ERROR is None, "status": status()}


# NB: not `/config` — Gradio's own client-discovery route lives there and shadowing it breaks `@gradio/client`.
@app.get("/studio-config")
def studio_config():
    """The canvas table and slider ranges, so the frontend never hardcodes a label the backend would reject."""
    import h3_lora

    state = getattr(PIPE.transformer, "_lora_state", None) if PIPE is not None else None
    sets = state["sets"] if state else {}
    return {
        "canvases": list(CANVASES),
        "default_canvas": DEFAULT_CANVAS,
        "min_duration": MIN_UI_DURATION,
        "max_duration": MAX_UI_DURATION,
        # The LoRA dropdown: value -> {label, suggested steps}.
        "loras": {
            **{
                name: {"label": spec["label"], "steps": {"larry": 6, "lightx": 4}.get(name, 6)}
                for name, spec in sets.items()
            },
            "off": {"label": "off (base model)", "steps": 28},
        },
        "default_lora": state["active"] if state else "off",
    }


@app.get("/", response_class=HTMLResponse)
def homepage():
    with open(os.path.join(os.path.dirname(os.path.abspath(__file__)), "index.html"), encoding="utf-8") as f:
        return f.read()


load_models()

if __name__ == "__main__":
    # allowed_paths: the /gradio_api/file= route only serves whitelisted directories.
    app.launch(show_error=True, allowed_paths=[OUTPUT_DIR])