VoxelModel-v1 / sample.py
TobiasLogic's picture
VoxelModel v1: tiny text-to-3D voxel diffusion
49ce4bd verified
Raw
History Blame Contribute Delete
2.49 kB
from __future__ import annotations
import argparse
import base64
import json
import os
import numpy as np
import torch
from safetensors.torch import load_file
from transformers import CLIPTextModel, CLIPTokenizer
from voxel_dit import VoxelDiT
MAX_TOKENS = 40
@torch.no_grad()
def sample(model, seq, pool, null_seq, null_pool, steps, cfg, dev):
B = seq.shape[0]
x = torch.randn(B, 1, 32, 32, 32, device=dev)
ns = null_seq.expand(B, -1, -1)
npool = null_pool.expand(B, -1)
dt = 1.0 / steps
for i in range(steps):
t = torch.full((B,), i * dt, device=dev)
with torch.autocast("cuda", dtype=torch.bfloat16):
vc = model(x, t, seq, pool)
vu = model(x, t, ns, npool)
x = x + (vu + cfg * (vc - vu)).float() * dt
return x
@torch.no_grad()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("prompts", nargs="+")
ap.add_argument("--out", default="samples.json")
ap.add_argument("--cfg", type=float, default=5.0)
ap.add_argument("--steps", type=int, default=50)
ap.add_argument("--threshold", type=float, default=0.0)
ap.add_argument("--device", default="cuda")
ap.add_argument("--weights", default="/root/runs/voxel_v1/model.safetensors")
ap.add_argument("--clip", default="openai/clip-vit-base-patch32")
args = ap.parse_args()
dev = args.device
model = VoxelDiT().to(dev).eval()
model.load_state_dict(load_file(args.weights))
tok = CLIPTokenizer.from_pretrained(args.clip)
txt = CLIPTextModel.from_pretrained(args.clip).to(dev).eval()
def enc(strings):
t = tok(strings, padding="max_length", max_length=MAX_TOKENS,
truncation=True, return_tensors="pt").to(dev)
o = txt(**t)
return o.last_hidden_state.float(), o.pooler_output.float()
seq, pool = enc(args.prompts)
null_seq, null_pool = enc([""])
x = sample(model, seq, pool, null_seq, null_pool, args.steps, args.cfg, dev)
grids = (x[:, 0] > args.threshold).to(torch.uint8).cpu().numpy()
out = []
for p, g in zip(args.prompts, grids):
occ = float(g.mean()) * 100
out.append({"prompt": p, "occS": round(occ, 1),
"s": base64.b64encode(np.packbits(g.reshape(-1)).tobytes()).decode()})
print(f" {occ:5.1f}% {p}")
json.dump(out, open(args.out, "w"))
print(f"[sample] {len(out)} grids -> {args.out} (cfg {args.cfg}, {args.steps} steps)")
if __name__ == "__main__":
main()