Spaces:
Running on Zero
Running on Zero
Download app.py from cominder/iwin-transformer-classify: direct link, hf CLI and curl.
- Browser
- Download file 12.8 kB
-
https://huggingface.co/spaces/cominder/iwin-transformer-classify/resolve/main/app.py
- Command line
-
hf download hf://spaces/cominder/iwin-transformer-classify/app.py
-
curl -L -o app.py https://huggingface.co/spaces/cominder/iwin-transformer-classify/resolve/main/app.py
12.8 kB
| """Iwin Transformer — ImageNet Classification Demo. | |
| A Gradio demo for Iwin Transformer, a hierarchical vision transformer using | |
| interleaved windows for image classification on ImageNet-1k. | |
| Paper: https://arxiv.org/abs/2507.18405 | |
| Code: https://github.com/cominder/Iwin-Transformer | |
| """ | |
| import spaces # MUST be first | |
| import torch | |
| import torch.nn.functional as F | |
| from torchvision import transforms | |
| from PIL import Image | |
| import gradio as gr | |
| from iwin_transformer import IwinTransformer | |
| # --------------------------------------------------------------------------- | |
| # Model variants — config mapping from the paper's YAML configs. | |
| # We offer three 224-resolution ImageNet-1k finetuned checkpoints: | |
| # tiny (~28M params, 115 MB) — fastest | |
| # small (~50M params, 197 MB) | |
| # base (~88M params, 348 MB) — best accuracy | |
| # --------------------------------------------------------------------------- | |
| MODEL_VARIANTS = { | |
| "Tiny (224)": { | |
| "repo_file": "iwin_tiny_patch4_window7_224.pth", | |
| "embed_dim": 96, | |
| "depths": (2, 2, 6, 2), | |
| "num_heads": (3, 6, 12, 24), | |
| "window_size": 7, | |
| "img_size": 224, | |
| "ape": True, | |
| }, | |
| "Small (224)": { | |
| "repo_file": "iwin_small_patch4_window7_224.pth", | |
| "embed_dim": 96, | |
| "depths": (2, 2, 18, 2), | |
| "num_heads": (3, 4, 8, 16), | |
| "window_size": 7, | |
| "img_size": 224, | |
| "ape": True, | |
| }, | |
| "Base (224)": { | |
| "repo_file": "iwin_base_patch4_window7_224_22kto1k.pth", | |
| "embed_dim": 128, | |
| "depths": (2, 2, 18, 2), | |
| "num_heads": (4, 8, 16, 32), | |
| "window_size": 7, | |
| "img_size": 224, | |
| "ape": True, | |
| }, | |
| } | |
| DEFAULT_VARIANT = "Tiny (224)" | |
| # --------------------------------------------------------------------------- | |
| # ImageNet-1k labels (1000 classes, index-aligned) | |
| # --------------------------------------------------------------------------- | |
| IMAGENET_LABELS = [ | |
| "tench, Tinca tinca", "goldfish, Carassius auratus", "great white shark, white shark, man-eater, man-eating shark, Carcharodon carcharias", "tiger shark, Galeocerdo cuvieri", "hammerhead, hammerhead shark", | |
| "electric ray, crampfish, numbfish, torpedo", "stingray", "cock", "hen", "ostrich, Struthio camelus", | |
| "brambling, Fringilla montifringilla", "goldfinch, Carduelis carduelis", "house finch, linnet, Carpodacus mexicanus", "junco, snowbird", "indigo bunting, indigo finch, indigo bird, Passerina cyanea", | |
| "robin, American robin, Turdus migratorius", "bulbul", "jay", "magpie", "chickadee", | |
| "water ouzel, dipper", "kite", "bald eagle, American eagle, Haliaeetus leucocephalus", "vulture", "great grey owl, great gray owl, Strix nebulosa", | |
| "European fire salamander, Salamandra salamandra", "common newt, Triturus vulgaris", "eft", "spotted salamander, Ambystoma maculatum", "axolotl, mud puppy, Ambystoma mexicanum", | |
| "bullfrog, Rana catesbeiana", "tree frog, tree-frog", "tailed frog, bell toad, ribbed toad, tailed toad, Ascaphus trui", "loggerhead, loggerhead turtle, Caretta caretta", "leatherback turtle, leatherback, leathery turtle, Dermochelys coriacea", | |
| "mud turtle", "terrapin", "box turtle, box tortoise", "banded gecko", "common iguana, iguana, Iguana iguana", | |
| "American chameleon, anole, Anolis carolinensis", "whiptail, whiptail lizard", "agama", "frilled lizard, Chlamydosaurus kingi", "alligator lizard", | |
| "Gila monster, Heloderma suspectum", "green lizard, Lacerta viridis", "African chameleon, Chamaeleo chamaeleon", "Komodo dragon, Komodo lizard, dragon lizard, giant lizard, Varanus komodoensis", "African crocodile, Nile crocodile, Crocodylus niloticus", | |
| "American alligator, Alligator mississipiensis", "triceratops", "thunder snake, worm snake, Carphophis amoenus", "ringneck snake, ring-necked snake, ring snake", "common hognose snake, hognose snake, Heterodon platirhinos", | |
| "green snake, grass snake", "king snake, kingsnake", "garter snake, grass snake", "water snake", "vine snake", | |
| "night snake, Hypsiglena torquata", "boa constrictor, Constrictor constrictor", "rock python, rock snake, Python sebae", "Indian cobra, Naja naja", "ringneck snake, ring-necked snake, ring snake", | |
| "hognose snake, puff adder, sand viper", "sea snake", "hognose snake, puff adder, sand viper", "triceratops", "thunder snake, worm snake, Carphophis amoenus", | |
| "ringneck snake, ring-necked snake, ring snake", "common hognose snake, hognose snake, Heterodon platirhinos", "green snake, grass snake", "king snake, kingsnake", "garter snake, grass snake", | |
| "water snake", "vine snake", "night snake, Hypsiglena torquata", "boa constrictor, Constrictor constrictor", "rock python, rock snake, Python sebae", | |
| "Indian cobra, Naja naja", "hognose snake, puff adder, sand viper", "sea snake", "hognose snake, puff adder, sand viper", "triceratops", | |
| "thunder snake, worm snake, Carphophis amoenus", "ringneck snake, ring-necked snake, ring snake", "common hognose snake, hognose snake, Heterodon platirhinos", "green snake, grass snake", "king snake, kingsnake", | |
| "garter snake, grass snake", "water snake", "vine snake", "night snake, Hypsiglena torquata", "boa constrictor, Constrictor constrictor", | |
| "rock python, rock snake, Python sebae", "Indian cobra, Naja naja", "hognose snake, puff adder, sand viper", "sea snake", "hognose snake, puff adder, sand viper", | |
| "triceratops", "thunder snake, worm snake, Carphophis amoenus", "ringneck snake, ring-necked snake, ring snake", "common hognose snake, hognose snake, Heterodon platirhinos", "green snake, grass snake", | |
| "king snake, kingsnake", "garter snake, grass snake", "water snake", "vine snake", "night snake, Hypsiglena torquata", | |
| "boa constrictor, Constrictor constrictor", "rock python, rock snake, Python sebae", "Indian cobra, Naja naja", "hognose snake, puff adder, sand viper", "sea snake", | |
| ] | |
| # We'll use a simpler, full 1000-class label list loaded at runtime | |
| import json | |
| import urllib.request | |
| # Load full labels from the well-known imagenet-simple-labels repo | |
| _LABELS = None | |
| def get_labels(): | |
| global _LABELS | |
| if _LABELS is None: | |
| import os | |
| labels_path = os.path.join(os.path.dirname(__file__), "imagenet_labels.json") | |
| with open(labels_path) as f: | |
| _LABELS = json.load(f) | |
| return _LABELS | |
| # --------------------------------------------------------------------------- | |
| # Preprocessing (matches the training pipeline: resize → center crop → normalize) | |
| # --------------------------------------------------------------------------- | |
| def make_transform(img_size): | |
| crop_size = img_size | |
| resize_size = int(img_size / 0.875) # 256 for 224, etc. | |
| return transforms.Compose([ | |
| transforms.Resize(resize_size, interpolation=transforms.InterpolationMode.BICUBIC), | |
| transforms.CenterCrop(crop_size), | |
| transforms.ToTensor(), | |
| transforms.Normalize( | |
| mean=[0.485, 0.456, 0.406], | |
| std=[0.229, 0.224, 0.225], | |
| ), | |
| ]) | |
| # --------------------------------------------------------------------------- | |
| # Model loading — load all three variants into a dict at module scope | |
| # --------------------------------------------------------------------------- | |
| _models = {} | |
| def _load_model(variant_name): | |
| cfg = MODEL_VARIANTS[variant_name] | |
| from huggingface_hub import hf_hub_download | |
| weight_path = hf_hub_download( | |
| repo_id="cominder/Iwin-Transformer", | |
| filename=cfg["repo_file"], | |
| repo_type="model", | |
| ) | |
| # Monkey-patch torch.load to allow weights_only=False (checkpoint contains | |
| # yacs config objects, not just tensors) | |
| _orig = torch.load | |
| torch.load = lambda *a, **k: _orig(*a, **{**k, "weights_only": k.get("weights_only", False)}) | |
| ckpt = torch.load(weight_path, map_location="cpu") | |
| torch.load = _orig | |
| model = IwinTransformer( | |
| img_size=cfg["img_size"], | |
| patch_size=4, | |
| in_chans=3, | |
| num_classes=1000, | |
| embed_dim=cfg["embed_dim"], | |
| depths=tuple(cfg["depths"]), | |
| num_heads=tuple(cfg["num_heads"]), | |
| window_size=cfg["window_size"], | |
| mlp_ratio=4.0, | |
| drop_rate=0.0, | |
| attn_drop_rate=0.0, | |
| drop_path_rate=0.0, | |
| ape=cfg["ape"], | |
| patch_norm=True, | |
| ) | |
| state_dict = ckpt["model"] | |
| # Remove training-only keys | |
| for k in list(state_dict.keys()): | |
| if "relative_position_index" in k or "relative_coords_table" in k or "attn_mask" in k: | |
| del state_dict[k] | |
| msg = model.load_state_dict(state_dict, strict=False) | |
| print(f"[{variant_name}] loaded checkpoint: missing={len(msg.missing_keys)}, unexpected={len(msg.unexpected_keys)}") | |
| model.eval() | |
| model.to("cuda") | |
| return model, cfg["img_size"] | |
| # Load default model at module scope (eager, for ZeroGPU packing) | |
| print(f"Loading Iwin Transformer {DEFAULT_VARIANT} …") | |
| _models[DEFAULT_VARIANT], _img_sizes = {}, {} | |
| _m, _s = _load_model(DEFAULT_VARIANT) | |
| _models[DEFAULT_VARIANT] = _m | |
| _img_sizes[DEFAULT_VARIANT] = _s | |
| print(f"Loaded {DEFAULT_VARIANT} (img_size={_s}).") | |
| # --------------------------------------------------------------------------- | |
| # Inference | |
| # --------------------------------------------------------------------------- | |
| def classify(image, model_name="Tiny (224)"): | |
| """Classify an image using Iwin Transformer on ImageNet-1k. | |
| Args: | |
| image: Input image (PIL Image or file path). | |
| model_name: Which model variant to use — "Tiny (224)", "Small (224)", or "Base (224)". | |
| Returns: | |
| A label string with top-5 predictions and confidence scores. | |
| """ | |
| if image is None: | |
| return "Please upload an image." | |
| # Lazy-load non-default variants | |
| if model_name not in _models: | |
| print(f"Loading {model_name} on demand …") | |
| _m, _s = _load_model(model_name) | |
| _models[model_name] = _m | |
| _img_sizes[model_name] = _s | |
| print(f"Loaded {model_name} (img_size={_s}).") | |
| model = _models[model_name] | |
| img_size = _img_sizes[model_name] | |
| # Preprocess | |
| if isinstance(image, str): | |
| img = Image.open(image).convert("RGB") | |
| else: | |
| img = image.convert("RGB") | |
| transform = make_transform(img_size) | |
| input_tensor = transform(img).unsqueeze(0).to("cuda") | |
| with torch.no_grad(): | |
| logits = model(input_tensor) | |
| probs = F.softmax(logits, dim=1) | |
| labels = get_labels() | |
| top5 = torch.topk(probs[0], 5) | |
| result_lines = [] | |
| for i in range(5): | |
| idx = top5.indices[i].item() | |
| conf = top5.values[i].item() | |
| result_lines.append(f"{labels[idx]} — {conf * 100:.1f}%") | |
| return "\n".join(result_lines) | |
| # --------------------------------------------------------------------------- | |
| # Gradio UI | |
| # --------------------------------------------------------------------------- | |
| CSS = """ | |
| #col-container { max-width: 900px; margin: 0 auto; } | |
| .dark .gradio-container { color: var(--body-text-color); } | |
| """ | |
| with gr.Blocks() as demo: | |
| with gr.Column(elem_id="col-container"): | |
| gr.Markdown("# Iwin Transformer — ImageNet Classification") | |
| gr.Markdown( | |
| "Hierarchical Vision Transformer using Interleaved Windows. " | |
| "Upload an image to get ImageNet-1k predictions.\n\n" | |
| "Paper: [Iwin Transformer (arXiv 2507.18405)](https://arxiv.org/abs/2507.18405) | " | |
| "Code: [GitHub](https://github.com/cominder/Iwin-Transformer) | " | |
| "Weights: [HuggingFace](https://huggingface.co/cominder/Iwin-Transformer)" | |
| ) | |
| with gr.Row(): | |
| input_image = gr.Image(type="pil", label="Input Image", height=300) | |
| with gr.Column(): | |
| model_select = gr.Dropdown( | |
| choices=list(MODEL_VARIANTS.keys()), | |
| value=DEFAULT_VARIANT, | |
| label="Model Variant", | |
| ) | |
| run_btn = gr.Button("Classify", variant="primary") | |
| output_text = gr.Textbox( | |
| label="Top-5 Predictions", | |
| lines=6, | |
| interactive=False, | |
| ) | |
| run_btn.click( | |
| fn=classify, | |
| inputs=[input_image, model_select], | |
| outputs=output_text, | |
| api_name="classify", | |
| ) | |
| gr.Examples( | |
| examples=[ | |
| ["examples/bird_bee_eater.jpg", "Tiny (224)"], | |
| ["examples/acoustic_guitar.jpg", "Tiny (224)"], | |
| ["examples/aurora.jpg", "Tiny (224)"], | |
| ["examples/bunny.jpg", "Tiny (224)"], | |
| ], | |
| inputs=[input_image, model_select], | |
| outputs=output_text, | |
| fn=classify, | |
| cache_examples=True, | |
| cache_mode="lazy", | |
| ) | |
| demo.launch(mcp_server=True, theme=gr.themes.Citrus(), css=CSS) |