File size: 17,625 Bytes
d661d4a
b8e6339
d661d4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
ff40e5a
d661d4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
fed6c68
d661d4a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
447
448
449
450
451
452
453
454
455
"""Bernini-Diffusers-v2 — reference-to-video (subject-to-video) demo.

Bernini couples a Qwen2.5-VL planner (which reads the reference images and the
instruction, then *plans* a target visual embedding with a flow-matching head)
to a Wan2.2-A14B MoE renderer (two 14B DiTs, high-noise + low-noise).

This Space mirrors the authors' own ``scripts/bernini_v2/run_r2v.sh`` /
``gradio_demo.py`` single-GPU path 1:1 (same guidance mode, omegas, planning
steps, system prompt and negative prompt); only the frame count / step count
defaults are lowered so a generation fits inside a ZeroGPU slot.
"""

import os

os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
os.environ.setdefault("HF_HUB_ENABLE_HF_TRANSFER", "1")

import spaces  # noqa: E402  (must precede torch / CUDA touching imports)

import gc  # noqa: E402
import logging  # noqa: E402
import random  # noqa: E402
import tempfile  # noqa: E402
import time  # noqa: E402

import gradio as gr  # noqa: E402
import torch  # noqa: E402
from huggingface_hub import hf_hub_download, snapshot_download  # noqa: E402

logging.basicConfig(level=logging.INFO, format="[%(asctime)s] %(name)s: %(message)s")
logging.getLogger("bernini.pipeline").setLevel(logging.INFO)

MODEL_ID = "ByteDance/Bernini-Diffusers-v2"


def _stat(tag):
    import shutil

    du = shutil.disk_usage("/tmp")
    rss = 0
    try:
        with open("/proc/self/status") as f:
            for line in f:
                if line.startswith("VmRSS"):
                    rss = int(line.split()[1]) / 1e6
    except Exception:
        pass
    print(f"[stat] {tag}: rss={rss:.1f}GB disk_used={du.used / 1e9:.1f}GB "
          f"free={du.free / 1e9:.1f}GB", flush=True)


# ---------------------------------------------------------------- weights ---
# The released checkpoint is fp32: `bernini/` alone is 180 GB, which blows past
# the Space's 150 GB disk quota.  So only the small components are materialised
# up-front; the 38 big shards are streamed one at a time, cast to bf16 straight
# into a meta-initialised model, and deleted immediately.  bf16 is the dtype the
# reference pipeline computes in anyway (`BerniniPipeline.weight_dtype`), so
# nothing is lost.  Peak disk for the shard stream is one shard (~5 GB).
#
# `mllm/*.safetensors` is skipped too: config.json sets `scratch_mllm: true`, so
# the MLLM is built from config and filled from the `bernini/` shards.
MODEL_DIR = snapshot_download(
    MODEL_ID,
    allow_patterns=[
        "config.json",
        "transformer_config.json",
        "transformer_2_config.json",
        "scheduler/*",
        "vae/*",
        "t5_text_encoder/*",
        "t5_tokenizer/*",
        "mllm/*.json",
        "mllm/*.txt",
        "mllm/*.model",
    ],
    max_workers=8,
)
_stat("after small snapshot")

# ------------------------------------------------------------------ model ---
import json  # noqa: E402

from accelerate import init_empty_weights  # noqa: E402
from safetensors import safe_open  # noqa: E402

from bernini.models import BerniniConfig, BerniniModel  # noqa: E402
from bernini.pipeline import BerniniPipeline, _localize_bernini_config  # noqa: E402
from diffusers.models import AutoencoderKLWan  # noqa: E402
from transformers import AutoProcessor, AutoTokenizer  # noqa: E402

config = BerniniConfig.from_pretrained(
    MODEL_DIR,
    use_unipc=True,
    use_src_id_rotary_emb=True,
    interpolate_src_id=True,
    max_trained_src_id=5,
)
_localize_bernini_config(config, MODEL_DIR)
config.mllm_attn_implementation = "sdpa"

with init_empty_weights():
    model = BerniniModel(config)
model.eval()
model.requires_grad_(False)
_stat("after meta init")

_index_path = hf_hub_download(MODEL_ID, f"{config.bernini_ckpt_subfolder}/model.safetensors.index.json")
_weight_map = json.load(open(_index_path))["weight_map"]
_shards = sorted(set(_weight_map.values()))
_pending = set(_weight_map)

for _i, _shard in enumerate(_shards, 1):
    _p = hf_hub_download(MODEL_ID, f"{config.bernini_ckpt_subfolder}/{_shard}")
    _sd = {}
    with safe_open(_p, framework="pt", device="cpu") as _f:
        for _k in _f.keys():
            _t = _f.get_tensor(_k)
            _sd[_k] = _t.to(torch.bfloat16) if _t.is_floating_point() else _t
            del _t
    model.load_state_dict(_sd, strict=False, assign=True)
    _pending -= set(_sd)
    del _sd
    for _f2 in {os.path.realpath(_p), _p}:
        try:
            os.remove(_f2)
        except OSError:
            pass
    gc.collect()
    print(f"[load] shard {_i}/{len(_shards)} {_shard}", flush=True)

_stat("after shard stream")
_meta = [n for n, p in model.named_parameters() if p.device.type == "meta"]
if _meta:
    print(f"[load] WARNING {len(_meta)} params still on meta, e.g. {_meta[:8]}", flush=True)
if _pending:
    print(f"[load] WARNING {len(_pending)} checkpoint keys unconsumed, e.g. {sorted(_pending)[:8]}", flush=True)

# transformer_2 is loaded inside diff_dec_low and attached back before sampling
setattr(model.diff_dec, "transformer_2", model.diff_dec_low.transformer_2)

t5_tokenizer = AutoTokenizer.from_pretrained(
    config.t5_tokenizer_path, subfolder=config.t5_tokenizer_subfolder, trust_remote_code=True
)
vit_processor = AutoProcessor.from_pretrained(
    config.processor_config_path,
    subfolder=config.processor_subfolder,
    padding_side="right",
    trust_remote_code=True,
)
vae = AutoencoderKLWan.from_pretrained(
    config.vae_model_path, subfolder=config.vae_subfolder, torch_dtype=torch.float32
)
vae.eval()
vae.requires_grad_(False)

PIPE = BerniniPipeline(config, model, vae, t5_tokenizer, vit_processor, "cuda")

# The two 14B renderer DiTs (~56 GB bf16) live on the GPU for the whole life of
# the Space.  The planner stack (MLLM / connector / vit head / T5 / VAE) is much
# smaller and the reference pipeline moves it on and off the device around its
# own phases, so it is left where that code expects to find it.
model.diff_dec.transformer.to("cuda")
model.diff_dec.transformer_2.to("cuda")
gc.collect()
_stat("after DiTs -> cuda")

# ------------------------------------------------------------------- task ---
# Verbatim from scripts/bernini_v2/run_r2v.sh
SYSTEM_PROMPT = "You are a helpful assistant specialized in subject-to-video generation."
NEG_PROMPT = (
    "vivid tones, overexposed, static, blurry details, subtitles, style, artwork, painting, "
    "image, motionless, overall grayish, worst quality, low quality, JPEG compression artifacts, "
    "ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn face, deformed, disfigured, "
    "malformed limbs, fused fingers, still frame, cluttered background, three legs, "
    "too many people in the background, walking backwards"
)
R2V = dict(
    guidance_mode="vae_txt_vit_wapg",
    max_image_size=842,
    flow_shift=5.0,
    fps=16,
    omega_txt=4.5,
    omega_tgt=1.5,
    omega_img=3.0,
    omega_vid=1.0,
    omega_scale=0.75,
    planning_step=50,
    vit_denoising_step=1,
    vit_txt_cfg=1.2,
    vit_img_cfg=1.0,
    eta=0.5,
    momentum=0.0,
    norm_threshold=(50.0, 50.0, 50.0),
)

RESOLUTIONS = {
    "Landscape · 848×480": (480, 848),
    "Portrait · 480×848": (848, 480),
    "Square · 640×640": (640, 640),
}
MAX_SEED = 2**31 - 1


def _coerce_gallery_paths(gallery_input):
    """gr.Gallery hands back a list of (path, caption) tuples."""
    if not gallery_input:
        return None
    out = []
    for item in gallery_input:
        if isinstance(item, (list, tuple)) and item:
            item = item[0]
        if isinstance(item, str):
            out.append(item)
        elif isinstance(item, dict) and item.get("path"):
            out.append(item["path"])
        elif hasattr(item, "name"):
            out.append(item.name)
    return out or None


def _estimate(*args, **kwargs):
    """Runtime scales with (denoising steps x latent tokens)."""
    try:
        n_images = max(1, len(args[0] or []))
        num_frames = int(args[2])
        steps = int(args[3])
        resolution = args[4]
    except Exception:
        return 420
    height, width = RESOLUTIONS.get(resolution, (480, 848))
    latent_frames = (int(num_frames) - 1) // 4 + 1
    tokens = latent_frames * (height // 16) * (width // 16)
    # Fitted on this Space (33f/848x480/16 steps unless noted):
    #   2 refs, 17f, 8 steps ->  95.1 s
    #   2 refs               -> 231.6 s warm / 254.3 s on a cold slot
    #   5 refs               -> 322.8 s
    # Planning cost scales with the reference count, sampling with steps x latent tokens.
    secs = 15.0 + 22.8 * n_images + 9.7e-4 * steps * tokens
    return int(min(800, max(90, secs * 1.15)))


@spaces.GPU(duration=_estimate, size="xlarge")
def generate(
    reference_images,
    prompt,
    num_frames=33,
    num_inference_steps=16,
    resolution="Landscape · 848×480",
    seed=42,
    randomize_seed=False,
    negative_prompt=NEG_PROMPT,
    omega_txt=4.5,
    omega_img=3.0,
    omega_tgt=1.5,
    omega_scale=0.75,
    progress=gr.Progress(track_tqdm=True),
):
    images = _coerce_gallery_paths(reference_images)
    if not images:
        raise gr.Error("Please add at least one reference image.")
    if len(images) > 8:
        raise gr.Error("Please use at most 8 reference images.")
    if not prompt or not prompt.strip():
        raise gr.Error("Please write a prompt describing the video you want.")

    if randomize_seed:
        seed = random.randint(0, MAX_SEED)
    height, width = RESOLUTIONS[resolution]

    out_path = tempfile.NamedTemporaryFile(suffix=".mp4", delete=False).name
    kwargs = dict(R2V)
    kwargs.update(
        omega_txt=float(omega_txt),
        omega_img=float(omega_img),
        omega_tgt=float(omega_tgt),
        omega_scale=float(omega_scale),
    )

    t0 = time.perf_counter()
    PIPE(
        "r2v",
        prompt.strip(),
        images=images,
        neg_prompt=negative_prompt or "",
        system_prompt=SYSTEM_PROMPT,
        num_frames=int(num_frames),
        height=int(height),
        width=int(width),
        num_inference_steps=int(num_inference_steps),
        seed=int(seed),
        output_path=out_path,
        **kwargs,
    )
    elapsed = time.perf_counter() - t0
    torch.cuda.empty_cache()
    print(f"[bernini] generated in {elapsed:.1f}s "
          f"({num_frames}f {width}x{height} {num_inference_steps} steps)", flush=True)
    return out_path, int(seed)


# --------------------------------------------------------------------- UI ---
EX1_PROMPT = (
    "The marble statue from image0, wearing the black T-shirt from image2, the tropical floral "
    "shorts from image3, and the pink cat-ear headphones from image1, sits on the wooden bench in "
    "the beach sunset setting from image4, facing the camera and gently bobbing and swaying to the "
    "music in a medium shot. Generate a video where the marble statue from image0 is the main "
    "subject, with the same muscular stone body, curly sculpted hair, and classical carved "
    "appearance, now humorously dressed in the black short-sleeve T-shirt from image2 with the "
    'white word "bernini" across the chest, the bright blue tropical floral shorts from image3 '
    "with large red, orange, and yellow flowers and green leaves, and the pink over-ear cat-ear "
    "headphones from image1. He is seated on the wooden bench from image4, centered in the frame "
    "and facing directly toward the camera in a medium shot. Keep the environment unchanged from "
    "image4: a seaside promenade with the wooden bench in the foreground, sandy beach and calm "
    "ocean behind it, palm trees rising on the left, and a vivid sunset sky glowing with warm "
    "orange, pink, and purple tones. He begins moving subtly and rhythmically as if listening to "
    "music through the headphones, gently nodding his head, swaying his upper body slightly, and "
    "rocking side to side in a natural music-driven motion, always remaining seated on the bench "
    "and facing the camera."
)
EX2_PROMPT = (
    "Place the male marble sculpture from image0 on the bench in image1, wearing the black T-shirt "
    'from image2 with the word "bernini" across the chest, holding the brown ceramic cup from '
    "image3 and slowly drinking from it with no steam visible, always facing the camera in a fixed "
    "medium shot. Keep the seaside sunset setting from image1 unchanged: the wooden bench centered "
    "on a paved path, palm trees on the left, and the beach, ocean and glowing sun in the "
    "background under a pink and orange sky. He starts seated upright holding the cup near his "
    "torso with a subtle rhythmic sway of the shoulders, then slowly lifts the cup toward his "
    "mouth in a controlled motion, gently tilts it and takes a sip, and finally lowers it while "
    "continuing a soft bobbing motion of the head and torso."
)

EXAMPLES = [
    [
        [
            "examples/source_img0.png",
            "examples/source_img1.png",
            "examples/source_img2.png",
            "examples/source_img3.png",
            "examples/source_img4.png",
        ],
        EX1_PROMPT,
    ],
    [
        [
            "examples/source_img0.png",
            "examples/source_img4.png",
            "examples/source_img2.png",
            "examples/source_img7.png",
        ],
        EX2_PROMPT,
    ],
]

CSS = """
#col-container { margin: 0 auto; max-width: 1100px; }
"""

with gr.Blocks(title="Bernini-Diffusers-v2") as demo:
    with gr.Column(elem_id="col-container"):
        gr.Markdown(
            """
# Bernini-Diffusers-v2 — reference-to-video

Drop in a few **reference images** (a subject, an outfit, a prop, a scene…), then describe the
video you want while pointing at them as `image0`, `image1`, … Bernini's Qwen2.5-VL planner reads
the references plus your instruction and plans a target visual embedding, which the Wan2.2-A14B
MoE renderer turns into a video.

[model](https://huggingface.co/ByteDance/Bernini-Diffusers-v2) ·
[code](https://github.com/bytedance/Bernini)
"""
        )

        with gr.Row():
            with gr.Column(scale=1):
                reference_images = gr.Gallery(
                    label="Reference images (order matters → image0, image1, …)",
                    file_types=["image"],
                    type="filepath",
                    columns=4,
                    height=240,
                    object_fit="contain",
                    interactive=True,
                    show_label=True,
                )
                prompt = gr.Textbox(
                    label="Prompt",
                    lines=6,
                    placeholder="The statue from image0, wearing the shirt from image1, sits on a "
                    "bench at sunset and gently sways to the music in a medium shot…",
                )
                run_btn = gr.Button("Generate video", variant="primary")
            with gr.Column(scale=1):
                video_out = gr.Video(label="Result", autoplay=True, height=380)
                used_seed = gr.Number(label="Seed used", interactive=False)

        with gr.Accordion("Advanced settings", open=False):
            with gr.Row():
                num_frames = gr.Slider(
                    label="Frames (16 fps)", minimum=17, maximum=49, step=4, value=33
                )
                num_inference_steps = gr.Slider(
                    label="Denoising steps", minimum=8, maximum=24, step=1, value=16
                )
            resolution = gr.Radio(
                label="Resolution",
                choices=list(RESOLUTIONS.keys()),
                value="Landscape · 848×480",
            )
            with gr.Row():
                seed = gr.Slider(label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=42)
                randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
            negative_prompt = gr.Textbox(label="Negative prompt", value=NEG_PROMPT, lines=3)
            gr.Markdown("Guidance weights — the defaults are the authors' `run_r2v.sh` values.")
            with gr.Row():
                omega_txt = gr.Slider(label="omega_txt", minimum=1.0, maximum=8.0, step=0.1, value=4.5)
                omega_img = gr.Slider(label="omega_img", minimum=0.0, maximum=8.0, step=0.1, value=3.0)
                omega_tgt = gr.Slider(label="omega_tgt", minimum=0.0, maximum=6.0, step=0.1, value=1.5)
                omega_scale = gr.Slider(label="omega_scale", minimum=0.0, maximum=1.0, step=0.05, value=0.75)

        gr.Markdown(
            "Longer clips and more steps look better but cost more GPU time. The defaults "
            "(33 frames ≈ 2 s at 16 fps, 16 steps) take about 4 minutes; the authors' reference "
            "setting is 81 frames / 40 steps, which does not fit in a single ZeroGPU slot."
        )

        gr.Examples(
            examples=EXAMPLES,
            inputs=[reference_images, prompt],
            outputs=[video_out, used_seed],
            fn=generate,
            cache_examples=True,
            cache_mode="lazy",
            label="Official Bernini r2v examples",
        )

    inputs = [
        reference_images,
        prompt,
        num_frames,
        num_inference_steps,
        resolution,
        seed,
        randomize_seed,
        negative_prompt,
        omega_txt,
        omega_img,
        omega_tgt,
        omega_scale,
    ]
    run_btn.click(fn=generate, inputs=inputs, outputs=[video_out, used_seed], api_name="generate")

demo.queue(max_size=12).launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)