File size: 16,880 Bytes
949ccb2
 
dd7a24f
949ccb2
 
 
 
 
 
2792700
949ccb2
dd7a24f
949ccb2
 
 
 
dd7a24f
949ccb2
dd7a24f
949ccb2
dd7a24f
 
949ccb2
dd7a24f
949ccb2
 
 
 
2792700
 
 
 
 
949ccb2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dd7a24f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
949ccb2
 
 
 
 
 
 
 
 
 
 
 
2792700
 
 
 
 
949ccb2
 
 
 
2792700
 
 
949ccb2
2792700
949ccb2
2792700
 
 
949ccb2
 
2792700
949ccb2
2792700
949ccb2
 
2792700
949ccb2
 
 
2792700
949ccb2
 
 
 
2792700
 
949ccb2
2792700
949ccb2
2792700
 
 
949ccb2
 
2792700
949ccb2
2792700
949ccb2
2792700
949ccb2
 
 
 
 
2792700
949ccb2
 
 
 
 
2792700
 
 
 
949ccb2
 
 
2792700
949ccb2
 
2792700
949ccb2
 
 
 
2792700
 
 
 
 
 
949ccb2
 
2792700
949ccb2
 
 
 
2792700
949ccb2
 
dd7a24f
 
 
 
2792700
dd7a24f
2792700
dd7a24f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2792700
dd7a24f
 
 
 
 
2792700
 
 
dd7a24f
2792700
 
dd7a24f
2792700
 
 
 
dd7a24f
 
2792700
dd7a24f
 
 
 
 
 
 
 
 
949ccb2
 
 
 
 
 
 
 
 
 
 
 
2792700
 
 
949ccb2
 
2792700
949ccb2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dd7a24f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
949ccb2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
dd7a24f
 
 
949ccb2
 
 
2792700
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
from __future__ import annotations

import json
import tempfile
import os

os.environ.setdefault("HF_HOME", r"D:\hf-cache")
os.environ.setdefault("HUGGINGFACE_HUB_CACHE", r"D:\hf-cache\hub")

import spaces
import gradio as gr
import librosa
import numpy as np
import torch
import trimesh
from diffusers import AutoencoderKL
from diffusers.pipelines.deprecated.audio_diffusion.mel import Mel
from huggingface_hub import hf_hub_download
from PIL import Image
from safetensors.torch import load_file
from transformers import AutoTokenizer, CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5TokenizerFast
from transformers import ClapTextModelWithProjection

from audio_dit import AudioDiT
from pixel_dit import DiT
from voxel_dit import VoxelDiT
from mmdit import MMDiT

# All boot-time loading and cached "null" embeddings stay on CPU.
# ZeroGPU only grants a GPU for the duration of an @spaces.GPU-decorated
# call, and import-time code runs without one, so nothing here should
# touch "cuda" directly. Each sampling function below moves the specific
# models/tensors it needs onto the allocated GPU (or CPU fallback) itself.
DEV = "cpu"
SCALE = 0.18215
CLIP_ID = "openai/clip-vit-base-patch32"
MAX_TOKENS = 40

print("[boot] loading shared CLIP text encoder...")
tokenizer = CLIPTokenizer.from_pretrained(CLIP_ID)
text_encoder = CLIPTextModel.from_pretrained(CLIP_ID).to(DEV).eval()


@torch.no_grad()
def encode(strings: list[str]):
    t = tokenizer(strings, padding="max_length", max_length=MAX_TOKENS, truncation=True, return_tensors="pt").to(DEV)
    o = text_encoder(**t)
    return o.last_hidden_state.float(), o.pooler_output.float()


null_seq, null_pool = encode([""])

print("[boot] loading PixelModel v5...")
pm5_weights = hf_hub_download("bench-labs/PixelModel-v5", "model.safetensors")
pm5_state = load_file(pm5_weights)
pixel_model = DiT(dim=384, depth=12, heads=6).to(DEV).eval()
pixel_model.load_state_dict({k[len("dit."):]: v for k, v in pm5_state.items() if k.startswith("dit.")})
vae = AutoencoderKL.from_pretrained("stabilityai/sd-vae-ft-mse").to(DEV).eval()

print("[boot] loading VoxelModel v1...")
vm1_weights = hf_hub_download("bench-labs/VoxelModel-v1", "model.safetensors")
voxel_model = VoxelDiT().to(DEV).eval()
voxel_model.load_state_dict(load_file(vm1_weights))

print("[boot] loading AudioModel v1...")
audio_cfg = json.load(open(hf_hub_download("bench-labs/AudioModel-v1", "config.json")))
a_dit = audio_cfg["dit"]
audio_model = AudioDiT(x_res=a_dit["x_res"], y_res=a_dit["y_res"],
                       text_seq_dim=a_dit["text_seq_dim"], text_pool_dim=a_dit["text_pool_dim"]).to(DEV).eval()
audio_model.load_state_dict(load_file(hf_hub_download("bench-labs/AudioModel-v1", "model_best.safetensors")))
audio_mel = Mel(x_res=a_dit["x_res"], y_res=a_dit["y_res"],
                sample_rate=audio_cfg["mel"]["sample_rate"], n_fft=audio_cfg["mel"]["n_fft"],
                hop_length=audio_cfg["mel"]["hop_length"], top_db=audio_cfg["mel"]["top_db"])
# Fast mel->STFT pseudo-inverse (librosa's pre-0.10 behavior). librosa 0.10+ uses an
# NNLS/L-BFGS solver here that allocates ~2GB and takes minutes on CPU for this
# 384x256 spectrogram; the pinv gives the same Griffin-Lim output in about a second.
audio_mel_pinv = np.linalg.pinv(
    librosa.filters.mel(sr=audio_mel.sr, n_fft=audio_mel.n_fft, n_mels=audio_mel.y_res, dtype=np.float32)
)
# Text-only CLAP tower: runs the same text_model + text_projection the reference
# sample.py uses, but skips the unused HTSAT audio tower (~150M params of dead weight).
audio_tokenizer = AutoTokenizer.from_pretrained("laion/clap-htsat-unfused")
audio_clap = ClapTextModelWithProjection.from_pretrained("laion/clap-htsat-unfused").to(DEV).eval()
audio_null_seq = audio_null_pool = None

print("[boot] loading PixelModel v6...")
v6_weights = hf_hub_download("bench-labs/PixelModel-v6", "model.safetensors")
pixel_model_v6 = MMDiT().to(DEV).eval()
pixel_model_v6.load_state_dict(load_file(v6_weights), strict=False)
t5_tokenizer = T5TokenizerFast.from_pretrained("google/flan-t5-base")
t5_encoder = T5EncoderModel.from_pretrained("google/flan-t5-base").to(DEV).eval()
vae_v6 = AutoencoderKL.from_pretrained("madebyollin/sdxl-vae-fp16-fix").to(DEV).float().eval()
null_v6_seq, null_v6_mask, null_v6_pool = None, None, None

print("[boot] ready.")


def _gpu_device() -> str:
    return "cuda" if torch.cuda.is_available() else "cpu"


@spaces.GPU(duration=60)
@torch.no_grad()
def sample_image(prompt: str, steps: int, cfg: float, seed: int, progress=gr.Progress()):
    if not prompt.strip():
        raise gr.Error("Type a prompt first.")
    device = _gpu_device()
    pixel_model.to(device)
    vae.to(device)
    steps = int(steps)
    g = torch.Generator(device=device).manual_seed(int(seed))
    seq, pool = encode([prompt])
    seq, pool = seq.to(device), pool.to(device)
    nseq, npool = null_seq.to(device), null_pool.to(device)
    x = torch.randn(1, 4, 32, 32, device=device, generator=g)
    dt = 1.0 / steps
    for i in progress.tqdm(range(steps), desc="sampling"):
        t = torch.full((1,), i * dt, device=device)
        vc = pixel_model(x, t, seq, pool)
        vu = pixel_model(x, t, nseq, npool)
        x = x + (vu + cfg * (vc - vu)) * dt
    img = vae.decode((x / SCALE)).sample
    img = ((img.clamp(-1, 1) + 1) / 2).permute(0, 2, 3, 1).cpu().numpy()[0]
    return (img * 255).round().astype(np.uint8)


@spaces.GPU(duration=90)
@torch.no_grad()
def sample_voxel(prompt: str, steps: int, cfg: float, threshold: float, seed: int, progress=gr.Progress()):
    if not prompt.strip():
        raise gr.Error("Type a prompt first.")
    device = _gpu_device()
    voxel_model.to(device)
    steps = int(steps)
    g = torch.Generator(device=device).manual_seed(int(seed))
    seq, pool = encode([prompt])
    seq, pool = seq.to(device), pool.to(device)
    nseq, npool = null_seq.to(device), null_pool.to(device)
    x = torch.randn(1, 1, 32, 32, 32, device=device, generator=g)
    dt = 1.0 / steps
    for i in progress.tqdm(range(steps), desc="sampling"):
        t = torch.full((1,), i * dt, device=device)
        vc = voxel_model(x, t, seq, pool)
        vu = voxel_model(x, t, nseq, npool)
        x = x + (vu + cfg * (vc - vu)) * dt
    grid = (x[0, 0] > threshold).cpu().numpy()
    if not grid.any():
        raise gr.Error("Nothing came back above the occupancy threshold — try lowering it or re-rolling the seed.")
    return grid_to_glb(grid)


@spaces.GPU(duration=120)
@torch.no_grad()
def sample_image_v6(prompt: str, steps: int, cfg: float, seed: int, progress=gr.Progress()):
    global null_v6_seq, null_v6_mask, null_v6_pool
    if not prompt.strip():
        raise gr.Error("Type a prompt first.")
    device = _gpu_device()
    pixel_model_v6.to(device)
    t5_encoder.to(device)
    vae_v6.to(device)
    steps = int(steps)

    def encode_v6(strings):
        t = t5_tokenizer(strings, padding="max_length", max_length=32, truncation=True, return_tensors="pt").to(device)
        seq = t5_encoder(**t).last_hidden_state.float()
        _, pool = encode(strings)
        return seq, t["attention_mask"].float(), pool.to(device)

    seq, mask, pool = encode_v6([prompt])
    if null_v6_seq is None:
        null_v6_seq, null_v6_mask, null_v6_pool = encode_v6([""])
    else:
        null_v6_seq, null_v6_mask, null_v6_pool = (
            null_v6_seq.to(device), null_v6_mask.to(device), null_v6_pool.to(device)
        )
    g = torch.Generator(device=device).manual_seed(int(seed))
    x = torch.randn(1, 4, 32, 32, device=device, generator=g)
    dt = 1.0 / steps
    for i in progress.tqdm(range(steps), desc="sampling"):
        t = torch.full((1,), i * dt, device=device)
        vc = pixel_model_v6(x, t, seq, mask, pool)
        vu = pixel_model_v6(x, t, null_v6_seq, null_v6_mask, null_v6_pool)
        x = x + (vu + cfg * (vc - vu)) * dt
    img = vae_v6.decode(x / vae_v6.config.scaling_factor).sample
    return (((img.clamp(-1, 1) + 1) / 2).permute(0, 2, 3, 1).cpu().numpy()[0] * 255).round().astype(np.uint8)


AUDIO_MAX_TOKENS = 32


@torch.no_grad()
def encode_audio(strings, device="cpu"):
    t = audio_tokenizer(strings, padding="max_length", max_length=AUDIO_MAX_TOKENS,
                        truncation=True, return_tensors="pt").to(device)
    out = audio_clap(**t)
    return out.last_hidden_state.float(), out.text_embeds.float()


def spectrogram_to_audio(mel, img):
    """Invert a mel spectrogram image back to audio (Griffin-Lim, 32 iters).



    Same 0..255 -> dB mapping as diffusers' Mel.image_to_audio, with the fast

    pseudo-inverse mel->STFT step (see the boot comment on audio_mel_pinv).

    """
    log_S = (np.frombuffer(img.tobytes(), dtype="uint8").reshape((img.height, img.width)).astype(np.float32)
             * mel.top_db / 255 - mel.top_db)
    power = librosa.db_to_power(log_S)
    stft_mag = np.clip(audio_mel_pinv @ power, 0, None) ** 0.5  # power=2.0 -> magnitude
    audio = librosa.griffinlim(stft_mag, n_iter=mel.n_iter, hop_length=mel.hop_length,
                               n_fft=mel.n_fft, window="hann")
    peak = np.abs(audio).max()
    return audio if peak == 0 else audio / peak * 0.9


@spaces.GPU(duration=90)
@torch.no_grad()
def sample_audio(prompt: str, steps: int, cfg: float, seed: int, progress=gr.Progress()):
    global audio_null_seq, audio_null_pool
    if not prompt.strip():
        raise gr.Error("Type a prompt first.")
    device = _gpu_device()
    audio_model.to(device)
    audio_clap.to(device)
    steps = int(steps)
    g = torch.Generator(device=device).manual_seed(int(seed))
    seq, pool = encode_audio([prompt], device=device)
    if audio_null_seq is None:
        audio_null_seq, audio_null_pool = encode_audio([""], device=device)
    else:
        audio_null_seq, audio_null_pool = audio_null_seq.to(device), audio_null_pool.to(device)
    x = torch.randn(1, 1, audio_model.y_res, audio_model.x_res, device=device, generator=g)
    dt = 1.0 / steps
    for i in progress.tqdm(range(steps), desc="sampling"):
        t = torch.full((1,), i * dt, device=device)
        vc = audio_model(x, t, seq, pool)
        vu = audio_model(x, t, audio_null_seq, audio_null_pool)
        x = x + (vu + cfg * (vc - vu)) * dt
    row = x[0, 0].clamp(-1, 1).float().cpu().numpy()
    img = Image.fromarray(((row + 1) * 127.5 + 0.5).astype(np.uint8))
    audio = spectrogram_to_audio(audio_mel, img)
    return audio_mel.get_sample_rate(), audio


def grid_to_glb(grid: np.ndarray) -> str:
    voxel = trimesh.voxel.VoxelGrid(encoding=grid)
    mesh = voxel.as_boxes()
    mesh.visual.face_colors = [180, 180, 190, 255]
    path = tempfile.NamedTemporaryFile(suffix=".glb", delete=False).name
    mesh.export(path)
    return path


with gr.Blocks(title="BenchLabs Models") as demo:
    gr.Markdown(
        "# BenchLabs Models\n"
        "Three tiny diffusion models, running live with on-demand GPU (ZeroGPU). "
        "All are under 45M trained parameters — the whole model fits in a PNG image "
        "if you're curious — see the model pages linked below."
    )
    with gr.Tab("Text → Image (PixelModel v6)"):
        gr.Markdown("A larger MMDiT model conditioned by T5 and CLIP. 256x256 output.")
        with gr.Row():
            with gr.Column():
                v6_prompt = gr.Textbox(label="Prompt", placeholder="a red fox sitting in a snowy forest")
                v6_steps = gr.Slider(10, 100, value=50, step=5, label="Detail (sampling steps)")
                v6_cfg = gr.Slider(1.0, 10.0, value=3.0, step=0.5, label="Prompt strength (CFG)")
                v6_seed = gr.Number(value=0, precision=0, label="Seed")
                v6_btn = gr.Button("Generate image", variant="primary")
            with gr.Column():
                v6_out = gr.Image(label="Result", type="numpy")
        v6_btn.click(sample_image_v6, [v6_prompt, v6_steps, v6_cfg, v6_seed], v6_out)
        gr.Examples(
            [["a red fox sitting in a snowy forest", 50, 3.0, 0],
             ["a lighthouse on a cliff at sunset", 50, 3.0, 0],
             ["a city street at night with neon signs", 50, 3.0, 0]],
            [v6_prompt, v6_steps, v6_cfg, v6_seed],
        )

    with gr.Tab("Text → 3D (VoxelModel v1)"):
        gr.Markdown(
            "Good at bulky objects: chairs, tables, cars, mushrooms. "
            "Thin objects (swords, keys) don't survive 32³ voxelization, in the training "
            "data or the model, so expect a blob rather than a blade."
        )
        with gr.Row():
            with gr.Column():
                vox_prompt = gr.Textbox(label="Prompt", placeholder="a wooden chair")
                vox_steps = gr.Slider(10, 50, value=25, step=1, label="Detail (sampling steps)")
                vox_cfg = gr.Slider(1.0, 10.0, value=5.0, step=0.5, label="Prompt strength (CFG)")
                vox_thresh = gr.Slider(-1.0, 1.0, value=0.0, step=0.05, label="Occupancy threshold")
                vox_seed = gr.Number(value=0, precision=0, label="Seed")
                vox_btn = gr.Button("Generate 3D model", variant="primary")
            with gr.Column():
                vox_out = gr.Model3D(label="Result")
        vox_btn.click(sample_voxel, [vox_prompt, vox_steps, vox_cfg, vox_thresh, vox_seed], vox_out)
        gr.Examples(
            [["a wooden chair", 25, 5.0, 0.0, 0],
             ["a purple mushroom", 25, 5.0, 0.0, 0],
             ["a small boat", 25, 5.0, 0.0, 0]],
            [vox_prompt, vox_steps, vox_cfg, vox_thresh, vox_seed],
        )

    with gr.Tab("Text → Audio (AudioModel v1)"):
        gr.Markdown(
            "17.8 seconds of sound at 22 kHz from a text prompt. The same tiny DiT + "
            "rectified-flow recipe as the other tabs, applied to a mel spectrogram "
            "image instead of pixels. Conditioned by the CLAP text tower; audio comes "
            "back via Griffin-Lim (32 iterations), so expect lo-fi, slightly phasey "
            "sound — there's no learned vocoder in v1."
        )
        with gr.Row():
            with gr.Column():
                aud_prompt = gr.Textbox(label="Prompt", placeholder="a dog barking")
                aud_steps = gr.Slider(10, 100, value=50, step=5, label="Detail (sampling steps)")
                aud_cfg = gr.Slider(1.0, 10.0, value=4.0, step=0.5, label="Prompt strength (CFG)")
                aud_seed = gr.Number(value=0, precision=0, label="Seed")
                aud_btn = gr.Button("Generate audio", variant="primary")
            with gr.Column():
                aud_out = gr.Audio(label="Result")
        aud_btn.click(sample_audio, [aud_prompt, aud_steps, aud_cfg, aud_seed], aud_out)
        gr.Examples(
            [["a dog barking", 50, 4.0, 0],
             ["rain falling on a roof", 50, 4.0, 0],
             ["footsteps on gravel", 50, 4.0, 0]],
            [aud_prompt, aud_steps, aud_cfg, aud_seed],
        )

    with gr.Tab("Text → Image (PixelModel v5)"):
        gr.Markdown(
            "Good at material and light: food, landscapes, skies, interiors. "
            "Weak on faces, hands, and anything needing precise structure or text."
        )
        with gr.Row():
            with gr.Column():
                img_prompt = gr.Textbox(label="Prompt", placeholder="a bowl of ramen with a soft boiled egg")
                img_steps = gr.Slider(10, 50, value=25, step=1, label="Detail (sampling steps)")
                img_cfg = gr.Slider(1.0, 10.0, value=5.0, step=0.5, label="Prompt strength (CFG)")
                img_seed = gr.Number(value=0, precision=0, label="Seed")
                img_btn = gr.Button("Generate image", variant="primary")
            with gr.Column():
                img_out = gr.Image(label="Result", type="numpy")
        img_btn.click(sample_image, [img_prompt, img_steps, img_cfg, img_seed], img_out)
        gr.Examples(
            [["a bowl of ramen with a soft boiled egg", 25, 5.0, 0],
             ["a wet cobblestone street at night", 25, 5.0, 0],
             ["a library of wooden shelves", 25, 5.0, 0]],
            [img_prompt, img_steps, img_cfg, img_seed],
        )

    gr.Markdown(
        "Models: [PixelModel v5](https://huggingface.co/bench-labs/PixelModel-v5) · "
        "[PixelModel v6](https://huggingface.co/bench-labs/PixelModel-v6) · "
        "[VoxelModel v1](https://huggingface.co/bench-labs/VoxelModel-v1) · "
        "[AudioModel v1](https://huggingface.co/bench-labs/AudioModel-v1)"
    )

if __name__ == "__main__":
    demo.queue(max_size=20).launch(server_name="0.0.0.0")