visual-tree / app.py
ShesterG's picture
Add DreamSim (human-aligned) encoder: LoRA-tuned DINO ViT-B/16 patch tokens + CLS attention
403700d
Raw
History Blame Contribute Delete
9.91 kB
"""
DINO Hierarchies — upload an image, get its condensed tree (hover) + reveal animation.
Runs DINOv3 ViT-L/16, builds the same condensed tree we precomputed for the static site,
and injects it into the existing JS widgets (no server-side rendering of the video).
"""
import os, io, json, html, base64, sys
print("[app] starting import", flush=True)
import numpy as np, torch
try:
import gradio as gr
print("[app] gradio", gr.__version__, flush=True)
except ImportError:
gr = None # allows importing the pipeline without gradio (local testing)
from PIL import Image
print("[app] transformers import...", flush=True)
from scipy.cluster.hierarchy import linkage, to_tree
from scipy.spatial.distance import squareform
from transformers import AutoModel, ViTMAEModel, CLIPVisionModel
RES, MS = 224, 4
DEV = "cuda" if torch.cuda.is_available() else "cpu"
TOKEN = os.environ.get("HF_TOKEN")
IMNET = (torch.tensor([0.485, 0.456, 0.406]).view(3, 1, 1),
torch.tensor([0.229, 0.224, 0.225]).view(3, 1, 1))
CLIPN = (torch.tensor([0.48145466, 0.4578275, 0.40821073]).view(3, 1, 1),
torch.tensor([0.26862954, 0.26130258, 0.27577711]).view(3, 1, 1))
# kind: how to strip prefix tokens / un-shuffle; patch: 224/patch = grid side
MODELS = {
"DINOv3 ViT-L · 0.3B (fast)": dict(repo="facebook/dinov3-vitl16-pretrain-lvd1689m", kind="dinov3", patch=16, norm=IMNET),
"DINOv3 ViT-H+ · 0.84B (slower)": dict(repo="facebook/dinov3-vith16plus-pretrain-lvd1689m", kind="dinov3", patch=16, norm=IMNET),
"DINOv3 ViT-7B · 6.7B (needs GPU)": dict(repo="facebook/dinov3-vit7b16-pretrain-lvd1689m", kind="dinov3", patch=16, norm=IMNET),
"MAE ViT-L (pixel-reconstruction)": dict(repo="facebook/vit-mae-large", kind="mae", patch=16, norm=IMNET),
"CLIP ViT-L/14 (language-aligned)": dict(repo="openai/clip-vit-large-patch14", kind="clip", patch=14, norm=CLIPN),
"DreamSim (human-aligned)": dict(repo="dreamsim:dino_vitb16", kind="dreamsim", patch=16, norm=None),
}
DEFAULT_MODEL = next(iter(MODELS))
TREE_HTML = open("tree_widget.html").read()
REVEAL_HTML = open("reveal_widget.html").read()
# single-slot cache: keep only the most-recently-used model (avoids OOM when switching)
_cur = {"repo": None, "model": None, "pre": None}
def get_model(spec):
repo = spec["repo"]
if _cur["repo"] != repo:
_cur["model"] = _cur["pre"] = None
import gc; gc.collect()
print(f"[app] loading {repo}", flush=True)
if spec["kind"] == "dreamsim":
from dreamsim import dreamsim as _ds # lazy: heavy dep, only when selected
cache = os.environ.get("DREAMSIM_CACHE", "/tmp/dreamsim_cache")
os.makedirs(cache, exist_ok=True)
mdl, pre = _ds(pretrained=True, dreamsim_type="dino_vitb16", cache_dir=cache, device=DEV)
_cur["model"], _cur["pre"] = mdl.eval(), pre
else:
dt = torch.bfloat16 if "vit7b16" in repo else torch.float32 # 7B in bf16 to fit
cls = {"dinov3": AutoModel, "mae": ViTMAEModel, "clip": CLIPVisionModel}[spec["kind"]]
mdl = cls.from_pretrained(repo, dtype=dt, token=TOKEN, attn_implementation="eager")
if spec["kind"] == "mae":
mdl.config.mask_ratio = 0.0 # keep ALL patches (MAE masks 75% by default)
_cur["model"] = mdl.eval().to(DEV)
_cur["repo"] = repo
return _cur["model"]
def encode(model, spec, x):
"""-> (feat[P,C], att[P]) with patches in spatial row-major order."""
if spec["kind"] == "dreamsim": # LoRA-tuned DINO ViT-B/16 inside DreamSim
vit = model.extractor_list[0].model
with torch.no_grad():
tok = vit.get_intermediate_layers(x, 1)[0] # (1, 1+P, C)
a = vit.get_last_selfattention(x) # (1, heads, S, S)
return tok[0, 1:].float().cpu().numpy(), a[0, :, 0, 1:].mean(0).float().cpu().numpy()
with torch.no_grad():
o = model(x, output_attentions=True)
h, a = o.last_hidden_state, o.attentions[-1][0] # a: (heads, S, S)
if spec["kind"] == "mae":
ids = o.ids_restore[0] # MAE shuffles tokens -> un-shuffle
feat, att = h[0, 1:][ids], a[:, 0, 1:][:, ids].mean(0)
else:
g = RES // spec["patch"]
prefix = h.shape[1] - g * g # strip CLS (+ registers for DINOv3)
feat, att = h[0, prefix:], a[:, 0, prefix:].mean(0)
return feat.float().cpu().numpy(), att.float().cpu().numpy()
def build_tree(image, model_label):
spec = MODELS.get(model_label, MODELS[DEFAULT_MODEL])
model = get_model(spec)
mdt = next(model.parameters()).dtype
img = image.convert("RGB").resize((RES, RES), Image.BICUBIC)
if spec["kind"] == "dreamsim":
x = _cur["pre"](img).to(DEV) # DreamSim ships its own transform
else:
mean, std = spec["norm"]
arr = np.asarray(img).astype(np.float32) / 255
x = ((torch.from_numpy(arr).permute(2, 0, 1) - mean) / std).unsqueeze(0).to(DEV, mdt)
feat, att = encode(model, spec, x)
att_pct = 100.0 * att / att.sum()
P = feat.shape[0]; g = int(round(P ** 0.5))
fn = feat / (np.linalg.norm(feat, axis=1, keepdims=True) + 1e-8)
root = to_tree(linkage(squareform(1.0 - fn @ fn.T, checks=False), method="average"))
nodes = []
def add(cn, depth, parent):
idx = len(nodes)
nodes.append({"leaves": cn.pre_order(lambda v: v.id), "children": [], "depth": depth,
"split_tau": None, "parent": parent})
cur, fo = cn, []
while not cur.is_leaf():
l, r = cur.left, cur.right
if l.count >= MS and r.count >= MS: break
small, big = (r, l) if l.count >= r.count else (l, r); fo.append(small); cur = big
if not cur.is_leaf():
nodes[idx]["split_tau"] = round(1.0 - cur.dist, 3)
for ch in (cur.left, cur.right):
nodes[idx]["children"].append(add(ch, depth + 1, idx))
for f in fo:
fi = len(nodes); nodes.append({"leaves": f.pre_order(lambda v: v.id), "children": [],
"depth": depth + 1, "split_tau": None, "parent": idx})
nodes[idx]["children"].append(fi)
return idx
add(root, 0, -1)
xpos = {}; cnt = [0]
def setx(i):
ch = nodes[i]["children"]
if not ch: xpos[i] = cnt[0]; cnt[0] += 1
else:
for c in ch: setx(c)
xpos[i] = float(np.mean([xpos[c] for c in ch]))
setx(0)
nleaves = cnt[0]; maxdepth = max(n["depth"] for n in nodes)
tree = {"g": g, "nleaves": nleaves, "maxdepth": maxdepth,
"nodes": [{"id": i, "x": round(xpos[i], 3), "depth": n["depth"], "parent": n["parent"],
"leaf": len(n["children"]) == 0, "n": len(n["leaves"]),
"att": round(float(att_pct[n["leaves"]].sum()), 1), # % of total attention
"attm": round(float(att_pct[n["leaves"]].mean()), 4), # mean attention % per patch (reveal order)
"attmv": round(float(att[n["leaves"]].mean()), 5), # raw mean attention value per patch
"tau": n["split_tau"], "patches": n["leaves"]} for i, n in enumerate(nodes)]}
buf = io.BytesIO(); img.save(buf, "PNG")
return tree, base64.b64encode(buf.getvalue()).decode()
def iframe(widget_html, tree, img_b64, height):
inject = (f'<script>window.TREE={json.dumps(tree)};'
f'window.IMG_SRC="data:image/png;base64,{img_b64}";</script>')
doc = widget_html.replace("<body>", "<body>" + inject, 1)
return (f'<iframe srcdoc="{html.escape(doc, quote=True)}" '
f'style="width:100%;height:{height}px;border:none;background:#0f1117;border-radius:8px"></iframe>')
def process(image, model_label):
if image is None:
return "<p style='color:#999'>Upload an image first.</p>", ""
tree, img_b64 = build_tree(image, model_label)
return iframe(TREE_HTML, tree, img_b64, 660), iframe(REVEAL_HTML, tree, img_b64, 700)
if gr is not None:
with gr.Blocks(title="Visual Tree", theme=gr.themes.Base()) as demo:
gr.Markdown("# 🌳 Visual Tree\nUpload an image → get its **condensed tree** (hover a node) "
"and the **reveal animation** (best-first by CLS attention).")
with gr.Row():
inp = gr.Image(type="pil", label="Upload an image", height=300)
with gr.Column(scale=0):
model_sel = gr.Dropdown(choices=list(MODELS.keys()), value=DEFAULT_MODEL, label="Encoder")
btn = gr.Button("Build tree + reveal", variant="primary")
gr.Markdown("<small>Compare encoder families: **DINOv3** (self-supervised, semantic), "
"**MAE** (pixel-reconstruction — more texture/appearance driven), "
"**CLIP** (language-aligned — 16×16 grid, patch-14), "
"**DreamSim** (tuned on *human* similarity judgments).<br>"
"ViT-H+/7B are large — slow on the free CPU; **ViT-7B realistically needs a GPU Space**. "
"First use of each model downloads its weights (one-time; DreamSim is ~1 GB).</small>")
with gr.Tab("Condensed tree (hover)"):
out_tree = gr.HTML()
with gr.Tab("Reveal animation"):
out_rev = gr.HTML()
btn.click(process, [inp, model_sel], [out_tree, out_rev])
inp.upload(process, [inp, model_sel], [out_tree, out_rev])
print("[app] launching gradio", flush=True)
demo.queue().launch() # module-level: HF Spaces imports this file, so no __main__ guard