multimodalart's picture
multimodalart HF Staff
Upload app.py with huggingface_hub
665d1a6 verified
Raw
History Blame Contribute Delete
10.3 kB
"""ProgResViT — progressive-resolution / progressive-width adaptive ViT.
Interactive ImageNet-1K classification demo that exposes the paper's
input-adaptive routing: round 1 runs a narrow subnetwork on a low-resolution
image, and only uncertain images continue to round 2 at higher resolution and
wider width.
Paper: https://huggingface.co/papers/2609.03216
Code: https://github.com/ds-kiel/ProgResViT
"""
import json
import os
import time
import spaces # must precede torch
import torch
import gradio as gr
from PIL import Image
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
from timm.data.transforms_factory import create_transform
from timm.models import create_model
# ---------------------------------------------------------------------------
# Model registry
# ---------------------------------------------------------------------------
# GMACs come from the authors' measured sweeps (results/RESULTS.md in the
# upstream repo): the threshold=0 row is the full two-round cost, the
# threshold=10 row (every image exits after round 1) is the round-1 cost.
VARIANTS = {
"160 → 384 · KD (84.9% top-1)": {
"repo": "NCPS/progresvit-deit-s-160-384-kd-imagenet1k",
"sizes": (160, 384),
"gmacs": (0.615, 16.152),
"top1": (73.940, 84.894),
"amp": True,
"threshold": 0.226,
},
"160 → 384 (83.7% top-1)": {
"repo": "NCPS/progresvit-deit-s-160-384-imagenet1k",
"sizes": (160, 384),
"gmacs": (0.615, 16.152),
"top1": (70.616, 83.714),
"amp": False,
"threshold": 0.267,
},
"192 → 240 · KD (83.8% top-1)": {
"repo": "NCPS/progresvit-deit-s-192-240-kd-imagenet1k",
"sizes": (192, 240),
"gmacs": (0.912, 6.267),
"top1": (76.018, 83.794),
"amp": False,
"threshold": 0.209,
},
"192 → 240 (82.2% top-1)": {
"repo": "NCPS/progresvit-deit-s-192-240-imagenet1k",
"sizes": (192, 240),
"gmacs": (0.912, 6.267),
"top1": (73.238, 82.202),
"amp": False,
"threshold": 0.356,
},
}
DEFAULT_VARIANT = "160 → 384 · KD (84.9% top-1)"
DEFAULT_THRESHOLD = VARIANTS[DEFAULT_VARIANT]["threshold"]
PROGRESS_STAGES = (3, 6) # attention heads active in round 1 / round 2
CACHE_VERSION = 1
with open(os.path.join(os.path.dirname(__file__), "imagenet_classes.json")) as f:
_ID2LABEL = json.load(f)
IMAGENET_CLASSES = [_ID2LABEL[str(i)] for i in range(1000)]
MODELS = {}
TRANSFORMS = {}
CROPS = {}
for _name, _spec in VARIANTS.items():
_cfg = json.load(open(hf_hub_download(_spec["repo"], "config.json")))
_model = create_model(
"progresvit",
pretrained=False,
num_classes=_cfg["num_classes"],
**_cfg["model_args"],
)
_state = load_file(hf_hub_download(_spec["repo"], "model.safetensors"))
_model.load_state_dict(_state, strict=True)
_pc = _cfg["pretrained_cfg"]
TRANSFORMS[_name] = create_transform(
input_size=tuple(_pc["input_size"]),
is_training=False,
interpolation=_pc["interpolation"],
mean=tuple(_pc["mean"]),
std=tuple(_pc["std"]),
crop_pct=_pc["crop_pct"],
crop_mode=_pc["crop_mode"],
crop_border_pixels=0,
use_prefetcher=False,
)
CROPS[_name] = int(_pc["input_size"][-1])
MODELS[_name] = _model.eval().to("cuda")
print(f"loaded {_name} from {_spec['repo']} (eval crop {CROPS[_name]})", flush=True)
def _topk_dict(logits: torch.Tensor, k: int = 5) -> dict:
probs = logits.float().softmax(dim=-1)[0]
values, indices = probs.topk(k)
return {IMAGENET_CLASSES[int(i)]: float(v) for v, i in zip(values, indices)}
@spaces.GPU(duration=15)
def classify(
image: Image.Image,
variant: str = DEFAULT_VARIANT,
threshold: float = DEFAULT_THRESHOLD,
) -> tuple:
"""Classify an image with ProgResViT's progressive, input-adaptive rounds.
Args:
image: input photograph to classify against the 1000 ImageNet-1K classes.
variant: which ProgResViT DeiT-S checkpoint to use (resolution schedule
and whether it was trained with knowledge distillation).
threshold: routing threshold on the round-1 top-10 prediction entropy.
The image exits after the cheap first round when its entropy falls
below this value; higher values exit more images and save more
compute.
Returns:
A tuple of (final top-5 prediction, routing report in markdown,
round-1 top-5 prediction, round-2 top-5 prediction).
"""
if image is None:
raise gr.Error("Please provide an image.")
spec = VARIANTS[variant]
model = MODELS[variant]
sizes = spec["sizes"]
g1, g2 = spec["gmacs"]
x = TRANSFORMS[variant](image.convert("RGB")).unsqueeze(0).to("cuda")
started = time.perf_counter()
with torch.inference_mode():
if spec["amp"]:
ctx = torch.autocast("cuda", dtype=torch.bfloat16)
else:
ctx = torch.autocast("cuda", enabled=False)
with ctx:
tokens1, logits1 = model._forward_stage(
x, 0, None, PROGRESS_STAGES, sizes
)
_, logits2 = model._forward_stage(
x, 1, tokens1, PROGRESS_STAGES, sizes
)
entropy = float(model.entropy(logits1.float())[0, 0])
elapsed = time.perf_counter() - started
exited_early = entropy < threshold
final_logits = logits1 if exited_early else logits2
used_gmacs = g1 if exited_early else g2
saving = 100.0 * (1.0 - used_gmacs / g2)
round1 = _topk_dict(logits1)
round2 = _topk_dict(logits2)
final = _topk_dict(final_logits)
if exited_early:
decision = (
f"**Exited after round 1.** Entropy `{entropy:.3f}` is below the "
f"threshold `{threshold:.3f}`, so the {sizes[1]} px round was skipped."
)
else:
decision = (
f"**Continued to round 2.** Entropy `{entropy:.3f}` is at or above the "
f"threshold `{threshold:.3f}`, so round 1's tokens were recycled and "
f"refined at {sizes[1]} px."
)
report = f"""### Routing
{decision}
| | Round 1 | Round 2 | This image |
|---|---|---|---|
| Input resolution | {sizes[0]} px | {sizes[1]} px | **{sizes[0] if exited_early else sizes[1]} px** |
| Active attention heads | {PROGRESS_STAGES[0]} / 6 | {PROGRESS_STAGES[1]} / 6 | **{PROGRESS_STAGES[0] if exited_early else PROGRESS_STAGES[1]} / 6** |
| Cumulative GMACs | {g1:.3f} | {g2:.3f} | **{used_gmacs:.3f}** |
| ImageNet top-1 if always stopped here | {spec['top1'][0]:.2f}% | {spec['top1'][1]:.2f}% | — |
Compute saved versus always running both rounds: **{saving:.1f}%** · inference {elapsed * 1000:.0f} ms
"""
return final, report, round1, round2
CSS = """
#col-container { max-width: 1180px; margin: 0 auto; }
.dark .gradio-container { color: var(--body-text-color); }
"""
# Ordered so the first rows tell the story: `red_fox` stays uncertain after round 1
# (which calls it a kit fox) and gets corrected in round 2, while `acoustic_guitar`
# is confident enough to exit after the cheap first round.
EXAMPLES = [
["examples/red_fox.jpg"],
["examples/acoustic_guitar.jpg"],
["examples/husky_dog.jpg"],
["examples/pizza_board.jpg"],
["examples/bird_kingfisher.jpg"],
["examples/chameleon.jpg"],
["examples/hot_air_balloon.jpg"],
["examples/vintage_camera.jpg"],
["examples/library_interior.jpg"],
["examples/spiral_staircase.jpg"],
["examples/monstera_plant.jpg"],
]
with gr.Blocks() as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
"""# ProgResViT — adaptive-compute image classification
An input-adaptive Vision Transformer that classifies progressively: round 1 runs a
**narrow** subnetwork on a **low-resolution** image, and only images whose prediction is
still uncertain continue to round 2 at **higher resolution** with a **wider** subnetwork,
reusing the tokens produced in round 1.
[Paper](https://huggingface.co/papers/2609.03216) · [Code](https://github.com/ds-kiel/ProgResViT) · [Checkpoints](https://huggingface.co/NCPS)
"""
)
with gr.Row():
with gr.Column():
image = gr.Image(label="Image", type="pil", height=340)
run = gr.Button("Classify", variant="primary")
variant = gr.Dropdown(
label="Checkpoint",
choices=list(VARIANTS),
value=DEFAULT_VARIANT,
)
threshold = gr.Slider(
label="Routing threshold (round-1 entropy)",
minimum=0.0,
maximum=2.0,
step=0.001,
value=DEFAULT_THRESHOLD,
info="0 = always run both rounds · higher = exit more images early",
)
with gr.Column():
final_out = gr.Label(label="Prediction", num_top_classes=5)
report_out = gr.Markdown(label="Routing report")
with gr.Accordion("Round-by-round predictions", open=False):
with gr.Row():
round1_out = gr.Label(label="Round 1 (low-res, narrow)", num_top_classes=5)
round2_out = gr.Label(label="Round 2 (high-res, wide)", num_top_classes=5)
gr.Examples(
examples=EXAMPLES,
inputs=[image],
outputs=[final_out, report_out, round1_out, round2_out],
fn=classify,
cache_examples=True,
cache_mode="lazy",
examples_per_page=12,
)
def _sync_threshold(name: str) -> float:
"""Reset the routing threshold to the checkpoint's reported operating point."""
return VARIANTS[name]["threshold"]
variant.change(_sync_threshold, inputs=variant, outputs=threshold)
run.click(
classify,
inputs=[image, variant, threshold],
outputs=[final_out, report_out, round1_out, round2_out],
api_name="classify",
)
if __name__ == "__main__":
demo.launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)