mit_b4_in1k / README.md
IMvision12's picture
fix readme.md
56353c8 verified
|
Raw
History Blame Contribute Delete
3.64 kB
---
pipeline_tag: image-classification
license: other
base_model: nvidia/mit-b4
library_name: kerasformers
tags:
- keras
- kerasformers
- image-classification
- mit
- backbone
- arxiv:2105.15203
- pytorch
- jax
- tf
---
## ***See [our collection](https://huggingface.co/collections/kerasformers/mit-segformer-encoder-6a6e81367fda42bf79b426e8) for all versions of MiT.***
# Run MiT with Keras 3: JAX, PyTorch, or TensorFlow
[![GitHub](https://img.shields.io/badge/GitHub-KerasFormers-black?logo=github)](https://github.com/IMvision12/KerasFormers) [![Docs](https://img.shields.io/badge/Docs-MiT-blue)](https://imvision12.github.io/KerasFormers/classification_backbones/) [![Collection](https://img.shields.io/badge/HF-MiT%20collection-yellow)](https://huggingface.co/collections/kerasformers/mit-segformer-encoder-6a6e81367fda42bf79b426e8)
# kerasformers/mit_b4_in1k
Paper: [SegFormer: Simple and Efficient Design for Semantic Segmentation with Transformers (arXiv:2105.15203)](https://arxiv.org/abs/2105.15203) · [HF Papers](https://huggingface.co/papers/2105.15203)
MiT is the hierarchical Mix Transformer encoder from SegFormer, also usable for ImageNet classification. For full SegFormer segmentation heads, see the SegFormer collection.
For more details on the model, please go to the upstream [model card](https://huggingface.co/nvidia/mit-b4).
Pure-**Keras 3** conversion of [`nvidia/mit-b4`](https://huggingface.co/nvidia/mit-b4) for [kerasformers](https://github.com/IMvision12/KerasFormers). One implementation runs unmodified on **TensorFlow / Torch / JAX**.
This is an **image-classification / backbone** checkpoint (`MiTImageClassify` / `MiTModel`).
## ✨ Quick start
```python
import os
os.environ["KERAS_BACKEND"] = "torch" # or "jax" / "tensorflow"
from PIL import Image
import numpy as np
from kerasformers.models.mit import MiTImageClassify, MiTModel
model = MiTImageClassify.from_weights("kerasformers/mit_b4_in1k")
backbone = MiTModel.from_weights(
"kerasformers/mit_b4_in1k", as_backbone=True
)
image = Image.open("your_image.jpg").convert("RGB")
image = image.resize((224, 224))
x = np.asarray(image, dtype="float32")[None] # (1, H, W, 3)
print(model(x).shape) # (1, num_classes)
feats = backbone(x)
print(len(feats), [tuple(f.shape) for f in feats])
```
Load any MiT variant the same way with `from_weights("kerasformers/<variant>")`:
| Variant | Hub |
|---|---|
| `mit_b0_in1k` | [`kerasformers/mit_b0_in1k`](https://huggingface.co/kerasformers/mit_b0_in1k) |
| `mit_b1_in1k` | [`kerasformers/mit_b1_in1k`](https://huggingface.co/kerasformers/mit_b1_in1k) |
| `mit_b2_in1k` | [`kerasformers/mit_b2_in1k`](https://huggingface.co/kerasformers/mit_b2_in1k) |
| `mit_b3_in1k` | [`kerasformers/mit_b3_in1k`](https://huggingface.co/kerasformers/mit_b3_in1k) |
| `mit_b4_in1k` | [`kerasformers/mit_b4_in1k`](https://huggingface.co/kerasformers/mit_b4_in1k) |
| `mit_b5_in1k` | [`kerasformers/mit_b5_in1k`](https://huggingface.co/kerasformers/mit_b5_in1k) |
## Tips
- Set `KERAS_BACKEND` **before** importing Keras / kerasformers.
- `MiTImageClassify` returns class logits; `MiTModel` returns features (`as_backbone=True` for multi-scale stages).
- See [docs](https://imvision12.github.io/KerasFormers/classification_backbones/) and [Loading Weights](https://imvision12.github.io/KerasFormers/loading_weights/).
- Upstream / timm checkpoints: `MiTImageClassify.from_weights("hf:nvidia/mit-b4")`.
## Special Thanks
A huge thank you to the MiT authors and the timm / Hub communities for creating and releasing these models.
License: see YAML `license` (usually matches the upstream checkpoint).