File size: 6,678 Bytes
cc4c5e8
 
 
 
7c27268
 
cc4c5e8
 
5d84fcd
7c27268
cc4c5e8
 
ef38a36
cc4c5e8
 
ef38a36
cc4c5e8
 
 
 
ef38a36
cc4c5e8
 
 
5d84fcd
 
cc4c5e8
 
 
 
 
5d84fcd
 
cc4c5e8
5d84fcd
cc4c5e8
 
 
 
 
 
5d84fcd
cc4c5e8
 
 
 
5d84fcd
cc4c5e8
5d84fcd
 
 
cc4c5e8
 
 
 
 
 
 
5d84fcd
cc4c5e8
 
 
 
 
 
 
 
 
5d84fcd
 
cc4c5e8
 
 
 
 
 
 
5d84fcd
cc4c5e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d84fcd
cc4c5e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d84fcd
cc4c5e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5d84fcd
 
cfcb1ea
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
from __future__ import annotations

import tempfile

import gradio as gr
import numpy as np
import torch
import trimesh
from diffusers import AutoencoderKL
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from transformers import CLIPTextModel, CLIPTokenizer

from pixel_dit import DiT
from voxel_dit import VoxelDiT

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] ready.")


@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.")
    steps = int(steps)
    g = torch.Generator(device=DEV).manual_seed(int(seed))
    seq, pool = encode([prompt])
    x = torch.randn(1, 4, 32, 32, device=DEV, generator=g)
    dt = 1.0 / steps
    for i in progress.tqdm(range(steps), desc="sampling"):
        t = torch.full((1,), i * dt, device=DEV)
        vc = pixel_model(x, t, seq, pool)
        vu = pixel_model(x, t, null_seq, null_pool)
        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).numpy()[0]
    return (img * 255).round().astype(np.uint8)


@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.")
    steps = int(steps)
    g = torch.Generator(device=DEV).manual_seed(int(seed))
    seq, pool = encode([prompt])
    x = torch.randn(1, 1, 32, 32, 32, device=DEV, generator=g)
    dt = 1.0 / steps
    for i in progress.tqdm(range(steps), desc="sampling"):
        t = torch.full((1,), i * dt, device=DEV)
        vc = voxel_model(x, t, seq, pool)
        vu = voxel_model(x, t, null_seq, null_pool)
        x = x + (vu + cfg * (vc - vu)) * dt
    grid = (x[0, 0] > threshold).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)


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"
        "Two tiny diffusion models, running live on CPU, no GPU behind this Space. "
        "Both are under 45M trained parameters, so generation is slower than a hosted API "
        "but the whole model fits in a PNG image if you're curious — see the model pages linked below."
    )
    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],
        )

    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],
        )

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

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