Spaces:
Running on Zero
Running on Zero
Download app.py from ivanmikhnenkov/tinydit: direct link, hf CLI and curl.
- Browser
- Download file 5.62 kB
-
https://huggingface.co/spaces/ivanmikhnenkov/tinydit/resolve/main/app.py
- Command line
-
hf download hf://spaces/ivanmikhnenkov/tinydit/app.py
-
curl -L -o app.py https://huggingface.co/spaces/ivanmikhnenkov/tinydit/resolve/main/app.py
5.62 kB
| """Hugging Face Space for tinydit-256 (ZeroGPU): prompt -> image, with an optional view of the | |
| sampling trajectory. Models load once at startup; the GPU is attached per request by ZeroGPU.""" | |
| import os, json, io, base64, torch, gradio as gr | |
| try: | |
| import spaces; GPU = spaces.GPU | |
| except Exception: # local / CPU fallback | |
| def GPU(*a, **k): return (lambda f: f) if not callable(a[0] if a else None) else a[0] | |
| from huggingface_hub import hf_hub_download, snapshot_download | |
| from safetensors.torch import load_file | |
| import diffusers | |
| from tinydit.model import TinyDiT | |
| from tinydit import text as T, schedule | |
| REPO = "ivanmikhnenkov/tinydit-256"; TOK = os.environ.get("HF_TOKEN") | |
| cfg = json.load(open(hf_hub_download(REPO, "config.json"))); stats = json.load(open(hf_hub_download(REPO, "stats.json"))) | |
| model = TinyDiT(latent_ch=cfg["latent_ch"], ctx_dim=cfg["ctx_dim"], dim=cfg["dim"], depth=cfg["depth"], heads=cfg["heads"], | |
| n_registers=cfg["n_registers"], register_block=cfg["register_block"], n_null=cfg["n_null"]).eval() | |
| sd = load_file(hf_hub_download(REPO, "model.safetensors")); own = model.state_dict() | |
| model.load_state_dict({k: v.to(own[k].dtype) for k, v in sd.items()}) | |
| tok, enc = T.load(snapshot_download("google/flan-t5-base", allow_patterns=["*.json", "*.model", "model.safetensors"]), device="cpu", dtype=torch.float32) | |
| vae = diffusers.AutoencoderKLFlux2.from_pretrained(os.path.join(snapshot_download("black-forest-labs/FLUX.2-dev", allow_patterns=["vae/*"], token=TOK), "vae"), | |
| torch_dtype=torch.float32).eval() | |
| mean = torch.tensor(stats["mean"]).view(1, -1, 1, 1); std = torch.tensor(stats["std"]).view(1, -1, 1, 1) | |
| model.to("cuda"); enc.to("cuda"); vae.to("cuda"); mean, std = mean.cuda(), std.cuda() | |
| SHAPES = {"256 脳 256": (256, 256), "288 脳 224 (4:3)": (288, 224), "224 脳 288 (3:4)": (224, 288), "320 脳 208 (3:2)": (320, 208), "208 脳 320 (2:3)": (208, 320)} | |
| def to_img(z): | |
| x = vae.decode((z * std + mean).float()).sample.clamp(-1, 1) | |
| return ((x[0].permute(1, 2, 0) + 1) * 127.5).round().byte().cpu().numpy() | |
| def generate(prompt, shape, steps, cfgs, seed, show_traj): | |
| if not prompt.strip(): raise gr.Error("Write a prompt first.") | |
| W, H = SHAPES[shape]; steps = int(steps); dev = "cuda" | |
| ctx, msk = T.embed_mixed(tok, enc, [prompt], [len(prompt.split()) > 30], device=dev); ctx = ctx.float() | |
| nctx0, nmsk0 = T.embed(tok, enc, [""], T.MAX_SHORT, device=dev) | |
| L = ctx.shape[1]; nctx = torch.zeros(1, L, ctx.shape[2], device=dev); nctx[:, :nctx0.shape[1]] = nctx0.float() | |
| nmsk = torch.zeros(1, L, dtype=torch.bool, device=dev); nmsk[:, :nmsk0.shape[1]] = nmsk0 | |
| g = torch.Generator(device=dev).manual_seed(int(seed)); x = torch.randn(1, cfg["latent_ch"], H // 8, W // 8, device=dev, generator=g) | |
| ts = schedule.shift(steps, cfg["sampler"]["shift"], device=dev); traj = [] | |
| keep = set(int(round(i)) for i in torch.linspace(0, steps - 1, 6).tolist()) if show_traj else set() | |
| for i in range(steps): | |
| t = ts[i].expand(1); dt = ts[i + 1] - ts[i] | |
| with torch.autocast("cuda", dtype=torch.bfloat16): | |
| v = model(torch.cat([x, x]), torch.cat([t, t]), torch.cat([ctx, nctx]), torch.cat([msk, nmsk])).float() | |
| vc, vu = v.chunk(2); vg = vu + float(cfgs) * (vc - vu) | |
| if i in keep: traj.append((to_img(x + (1 - ts[i]) * vg), f"prediction at step {i+1}/{steps}, t={float(ts[i]):.2f}")) | |
| x = x + vg * dt | |
| return to_img(x), traj | |
| with gr.Blocks(title="tinydit-256") as demo: | |
| gr.Markdown("# tinydit-256\nA 210M text-to-image diffusion transformer trained from scratch on one GPU in 3.5 days. " | |
| "Defaults are the training-time sampler (20 steps, CFG 4). " | |
| "[Code](https://github.com/ivanmikhnenkov/tinydit) 路 [Weights](https://huggingface.co/ivanmikhnenkov/tinydit-256) 路 [ivanmikhnenkov.com](https://ivanmikhnenkov.com)") | |
| with gr.Row(): | |
| with gr.Column(scale=1): | |
| prompt = gr.Textbox(label="Prompt", value="a red tractor parked next to a blue rowing boat on a sandy beach", lines=3) | |
| shape = gr.Dropdown(list(SHAPES), value="256 脳 256", label="Size (the five training shapes)") | |
| steps = gr.Slider(4, 50, value=20, step=1, label="Steps"); cfgs = gr.Slider(1.0, 8.0, value=4.0, step=0.5, label="Guidance (CFG)") | |
| seed = gr.Number(value=0, precision=0, label="Seed"); show_traj = gr.Checkbox(value=True, label="Show how the image forms (6 intermediate predictions)") | |
| btn = gr.Button("Generate", variant="primary") | |
| with gr.Column(scale=1): | |
| out = gr.Image(label="Result", type="numpy") | |
| traj = gr.Gallery(label="The model's prediction of the final image along the sampling steps", columns=3, height=260) | |
| gr.Examples([["three green apples on a white plate next to a black coffee cup", "256 脳 256", 20, 4.0, 5, True], | |
| ["a lighthouse on a rocky coast at sunset", "320 脳 208 (3:2)", 20, 4.0, 1, True], | |
| ["a golden retriever wearing sunglasses sitting on a yellow armchair", "224 脳 288 (3:4)", 20, 4.0, 2, True], | |
| ["a bowl of ramen with a soft-boiled egg and green onions on a dark table", "256 脳 256", 20, 4.0, 2, True]], | |
| inputs=[prompt, shape, steps, cfgs, seed, show_traj]) | |
| btn.click(generate, [prompt, shape, steps, cfgs, seed, show_traj], [out, traj]) | |
| prompt.submit(generate, [prompt, shape, steps, cfgs, seed, show_traj], [out, traj]) | |
| demo.queue(max_size=16).launch() | |