tedzh17 commited on
Commit
7db838f
·
verified ·
1 Parent(s): 6794be9

Update README.md

Browse files
Files changed (1) hide show
  1. README.md +129 -95
README.md CHANGED
@@ -13,109 +13,129 @@ tags:
13
  - image-feature-extraction
14
  ---
15
 
16
- # STELLAR — Sparse Visual Representations via Spatial–Semantic Factorization
17
 
18
- **STELLAR** learns a **unified sparse visual representation** that supports both
19
- **reconstruction** and **semantics** using as few as **16 tokens**. By factorizing
20
- *"what"* (semantics) from *"where"* (spatial layout), each image is encoded as the
21
- low-rank product of a **localization** matrix and a **semantics** matrix.
 
 
 
 
 
 
 
22
 
23
  <p align="center">
24
  <img src="factorization.svg" alt="Spatial–semantic factorization" width="720">
25
  </p>
26
 
27
- - 📄 **Paper:** [arXiv:2602.01905](https://arxiv.org/abs/2602.01905) (ICML 2026)
28
- - 💻 **Code:** [github.com/microsoft/STELLAR](https://github.com/microsoft/STELLAR)
 
29
 
30
  These checkpoints contain the **full set of trained STELLAR modules** (encoder, sparse
31
  tokens, projections, reconstruction decoder, and clustering heads), so a single file
32
  supports feature extraction, image reconstruction, and continued pretraining. All
33
  models are self-supervised on **ImageNet-1K** at 224×224.
34
 
35
- ## Highlights
36
-
37
- - **Sparse & unified** — one small set of tokens serves both high-level semantics
38
- and pixel-level reconstruction.
39
- - **Factorized latents** — each token captures a concept (*what*) together with a
40
- spatial map of *where* it appears.
41
- - **Strong on both axes** — STELLAR-H reaches **2.60 FID** (reconstruction) and
42
- **79.1%** ImageNet linear-probing accuracy with just **16 tokens**.
43
 
44
  ## Available models
45
 
46
- | Model | Backbone | Tokens | Params | Type | File |
47
  | :--- | :--- | :---: | :---: | :--- | :--- |
48
- | `stellar-b16` | ViT-B/16 | 16 | 88M | main | [`stellar-b16.safetensors`](stellar-b16.safetensors) |
49
- | `stellar-l16` | ViT-L/16 | 16 | 307M | main | [`stellar-l16.safetensors`](stellar-l16.safetensors) |
50
- | `stellar-h16` | ViT-H/14 | 16 | 636M | main | [`stellar-h16.safetensors`](stellar-h16.safetensors) |
51
- | `stellar-b8` | ViT-B/16 | 8 | 88M | ablation | [`stellar-b8.safetensors`](stellar-b8.safetensors) |
52
- | `stellar-b24` | ViT-B/16 | 24 | 88M | ablation | [`stellar-b24.safetensors`](stellar-b24.safetensors) |
53
 
54
- The main models (`b16`, `l16`, `h16`) are recommended for downstream use; the 8- and
55
- 24-token base models are ablations on the number of sparse tokens.
 
56
 
57
  ## Usage
58
 
59
- Install the STELLAR code and the Hub helpers:
60
 
61
  ```bash
62
- pip install huggingface_hub safetensors
63
- git clone https://github.com/microsoft/STELLAR && cd STELLAR
64
- pip install -r requirements.txt
65
  ```
66
 
 
 
 
 
67
  ### Quick start
68
 
69
- From the STELLAR code directory, use the [`load_stellar.py`](load_stellar.py) helper
70
- (it downloads the weights from the Hub for you):
71
 
72
  ```python
73
  import torch
74
- from load_stellar import load_stellar, list_models
75
-
76
- print(list_models()) # ['stellar-b16', 'stellar-l16', ...]
77
- model = load_stellar("stellar-b16") # purpose="encode" (default)
78
 
79
- # RGB image in [0, 1], resized to 224×224 (ImageNet normalization is applied internally)
80
- image = torch.rand(1, 3, 224, 224)
 
 
81
  with torch.no_grad():
82
  out = model.encode(image)
83
 
84
- out["sparse"] # (1, K, D) sparse concept tokens ("what")
85
- out["spatial"] # (1, P, K) per-token spatial maps ("where")
86
- out["dense"] # (1, P, D) dense per-patch features
87
- out["cls"] # (1, 1, D) global image token
88
  ```
89
 
90
- ### Reconstruction & continued pretraining
 
 
 
 
 
 
 
 
91
 
92
- The same checkpoint can be loaded for other purposes via the `purpose` argument. Image
93
- reconstruction and continued pretraining use the decoder, which predicts
94
- [MaskGIT-VQGAN](https://huggingface.co/fun-research/TiTok) tokens — pass the tokenizer
95
- path as `vq_model`:
96
 
97
  ```python
98
- # 1. encode -> factorized features (sparse concept tokens + spatial maps)
99
- model = load_stellar("stellar-b16", purpose="reconstruct", vq_model=VQGAN_PATH)
100
- features = model.encode(image) # dict: sparse (B,K,D), spatial (B,P,K), ...
101
-
102
- # 2. decode the factorized features -> VQGAN decoder -> pixels
103
- out = model.reconstruct(features) # or model.reconstruct(features["sparse"], features["spatial"])
104
- pixels = out["reconstruction"] # (B, 3, H, W) RGB in [0, 1]
105
- # 224x224 for /16 models, 256x256 for the /14 H model
106
- # out["tokens"] : (B, P) predicted VQGAN token ids
107
- # out["logits"] : (B, P, 1024) raw codebook logits
108
-
109
- # continued pretraining (all modules, gradients enabled)
110
- model = load_stellar("stellar-b16", purpose="pretrain", vq_model=VQGAN_PATH)
111
- losses = model({"image": image, "labels": labels, ...})["predictions"]
112
- ```
113
 
114
- `reconstruct` is the **decoder half** of STELLAR: it takes the factorized features and
115
- runs low-rank dense map → ViT decoder → VQGAN decoder to return RGB pixels. See
116
- [`examples/reconstruction.ipynb`](examples/reconstruction.ipynb) for an end-to-end demo
117
- that loads an image and displays the reconstruction.
 
 
 
 
 
 
 
 
 
 
 
118
 
 
 
 
 
 
 
 
119
 
120
  ### What the model returns
121
 
@@ -130,39 +150,38 @@ that loads an image and displays the reconstruction.
130
  `B` = batch, `K` = number of sparse tokens, `P` = number of patches (196 for /16 at
131
  224², 256 for /14), `D` = embedding dim (768 / 1024 / 1280 for B / L / H).
132
 
133
- ### Loading the weights manually
134
 
135
- ```python
136
- import json, torch
137
- from huggingface_hub import hf_hub_download
138
- from safetensors.torch import load_file
139
- from src.models.stellar_model import STELLARModel
140
-
141
- repo = "microsoft/STELLAR"
142
- cfg = json.load(open(hf_hub_download(repo, "config.json")))["models"]["stellar-b16"]
143
- state = load_file(hf_hub_download(repo, cfg["weights"]))
144
-
145
- model = STELLARModel(
146
- num_sparse_tokens=cfg["num_sparse_tokens"],
147
- num_decoder_layers=cfg["num_decoder_layers"],
148
- spatial_temp=cfg["spatial_temp"],
149
- vit_pretrained=cfg["backbone"],
150
- do_recon=False, do_clustering=False, vq_model=None,
151
- )
152
- model.load_state_dict(state, strict=False) # encoder-only build ignores decoder/heads
153
- model.eval()
154
- features = model.encode(torch.rand(1, 3, 224, 224))
155
- ```
156
 
157
- > **Tip:** download with `huggingface_hub` (as above) rather than `git clone` so that
158
- > downloads are registered on the Hub — `git clone` is not counted in download stats.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
159
 
160
  ## Model details
161
 
162
  - **Architecture:** ViT encoder (MAE-initialized) + learned sparse latent queries with
163
  spatial–semantic factorization.
164
  - **Pretraining data:** ImageNet-1K (self-supervised; labels not used).
165
- - **Input:** RGB images in `[0, 1]`, resized to 224×224 (bicubic). ImageNet mean/std
166
  normalization is applied **inside** the model — pass raw `[0, 1]` images.
167
  - **Weights:** the complete set of trained STELLAR modules (encoder, sparse tokens,
168
  projections, reconstruction decoder, and clustering heads), stored in `safetensors`.
@@ -172,13 +191,28 @@ features = model.encode(torch.rand(1, 3, 224, 224))
172
 
173
  ## Intended uses & limitations
174
 
175
- - **Intended use:** extracting compact sparse/dense visual features for downstream
176
- recognition, segmentation, retrieval, reconstruction, and analysis.
177
- - **Limitations:** pretrained on ImageNet-1K at 224×224, so features reflect that
178
- distribution; performance on very different domains (e.g. medical, satellite) may
179
- require fine-tuning. The models are research artifacts and are not safety-tested for
 
 
 
 
180
  production decision-making.
181
 
 
 
 
 
 
 
 
 
 
 
 
182
  ## Citation
183
 
184
  ```bibtex
@@ -187,10 +221,10 @@ features = model.encode(torch.rand(1, 3, 224, 224))
187
  author = {Zhao, Theodore Zhengde and Kiblawi, Sid and Yang, Jianwei and Usuyama, Naoto and Tan, Reuben and Codella, Noel C and Naumann, Tristan and Poon, Hoifung and Wei, Mu},
188
  booktitle = {International Conference on Machine Learning (ICML)},
189
  year = {2026},
190
- url = {https://arxiv.org/abs/2602.01905},
191
  }
192
  ```
193
 
194
  ## License
195
 
196
- Released under the [MIT License](LICENSE).
 
13
  - image-feature-extraction
14
  ---
15
 
16
+ # STELLAR: Learning Sparse Visual Representations via Spatial–Semantic Factorization
17
 
18
+ **How many tokens are needed for one image?**
19
+
20
+ We show that with a **spatial–semantic factorized representation**, **16 semantic
21
+ tokens** paired with explicit spatial maps are enough for strong **image recognition
22
+ and reconstruction**. STELLAR separates **what** an image contains from **where** it
23
+ appears, representing the image as the low-rank product of a localization matrix
24
+ and a semantics matrix.
25
+
26
+ STELLAR-H achieves **2.60 reconstruction FID** and **79.10% ImageNet linear-probing
27
+ accuracy**, using approximately **90% fewer latent values** than a dense grid when
28
+ counting both factors.
29
 
30
  <p align="center">
31
  <img src="factorization.svg" alt="Spatial–semantic factorization" width="720">
32
  </p>
33
 
34
+ - **ICML 2026:** [Paper](https://openreview.net/pdf?id=ysOOfySED6) / [arXiv](https://arxiv.org/abs/2602.01905)
35
+ - **Code and configs:** [microsoft/STELLAR](https://github.com/microsoft/STELLAR)
36
+ - **Training and evaluation:** [Usage guide](https://github.com/microsoft/STELLAR/blob/main/docs/usage.md)
37
 
38
  These checkpoints contain the **full set of trained STELLAR modules** (encoder, sparse
39
  tokens, projections, reconstruction decoder, and clustering heads), so a single file
40
  supports feature extraction, image reconstruction, and continued pretraining. All
41
  models are self-supervised on **ImageNet-1K** at 224×224.
42
 
43
+ The headline results use the evaluation protocols in our ICML work. Reconstruction
44
+ FID is **not generation FID**; loading a checkpoint or running a new probe does not
45
+ automatically reproduce those numbers. The released files contain pretraining
46
+ weights, not the separately finetuned B/L reconstruction probes.
 
 
 
 
47
 
48
  ## Available models
49
 
50
+ | Model | Backbone | Semantic tokens | Feature width | Type | Weights |
51
  | :--- | :--- | :---: | :---: | :--- | :--- |
52
+ | `stellar-b16` | ViT-B/16 | 16 | 768 | main | [safetensors](stellar-b16.safetensors) |
53
+ | `stellar-l16` | ViT-L/16 | 16 | 1024 | main | [safetensors](stellar-l16.safetensors) |
54
+ | `stellar-h16` | ViT-H/14 | 16 | 1280 | main | [safetensors](stellar-h16.safetensors) |
55
+ | `stellar-b8` | ViT-B/16 | 8 | 768 | ablation | [safetensors](stellar-b8.safetensors) |
56
+ | `stellar-b24` | ViT-B/16 | 24 | 768 | ablation | [safetensors](stellar-b24.safetensors) |
57
 
58
+ Start with **B16** for a smaller backbone; use L16 or H16 when memory permits.
59
+ The B8 and B24 variants study the number of sparse tokens. The token count in the
60
+ model name is independent of the backbone patch size: H16 uses a ViT-H/14 encoder.
61
 
62
  ## Usage
63
 
64
+ Use Python 3.10 and run from the GitHub code directory:
65
 
66
  ```bash
67
+ git clone https://github.com/microsoft/STELLAR.git
68
+ cd STELLAR
69
+ python -m pip install -r requirements-inference.txt
70
  ```
71
 
72
+ Feature extraction needs no Azure account, Olympus trainer, VQGAN weights, or
73
+ separate MAE download. This is a custom PyTorch model; use the repository loader,
74
+ not `transformers.AutoModel` or `pipeline()`.
75
+
76
  ### Quick start
77
 
78
+ The [GitHub loader](https://github.com/microsoft/STELLAR/blob/main/load_stellar.py)
79
+ downloads and validates the required weights. Replace `your_image.jpg` with a local image:
80
 
81
  ```python
82
  import torch
83
+ from load_stellar import load_stellar
84
+ from examples.common import read_image
 
 
85
 
86
+ revision = "6794be9a20fb9d3944c5bcc6003512a347fa518f"
87
+ device = "cuda" if torch.cuda.is_available() else "cpu"
88
+ model = load_stellar("stellar-b16", revision=revision, device=device)
89
+ image = read_image("your_image.jpg").unsqueeze(0).to(device)
90
  with torch.no_grad():
91
  out = model.encode(image)
92
 
93
+ print(out["sparse"].shape)
94
+ print(out["spatial"].shape)
 
 
95
  ```
96
 
97
+ For B16, these shapes are `(1, 16, 768)` and `(1, 196, 16)`. The helper converts
98
+ to RGB, resizes the short side to 256 and center-crops to 224. Inputs must be
99
+ floating-point RGB tensors in `[0, 1]`; **do not apply ImageNet normalization**,
100
+ which is already performed inside the model.
101
+
102
+ Pin the Hub revision for repeatable downloads. After caching the weights,
103
+ `local_files_only=True` enables offline loading. Set `HF_HUB_CACHE` before starting
104
+ Python to choose a cache directory. Allow roughly 0.5 GB per B checkpoint, 1.4 GB
105
+ for L16, and 2.7 GB for H16, plus the tokenizer if reconstructing.
106
 
107
+ ### Reconstruction
108
+
109
+ Continuing from the quick start, download the external
110
+ [MaskGIT-VQGAN tokenizer](https://huggingface.co/fun-research/TiTok) and load the decoder:
111
 
112
  ```python
113
+ from huggingface_hub import hf_hub_download
114
+ from torchvision.transforms.functional import to_pil_image
 
 
 
 
 
 
 
 
 
 
 
 
 
115
 
116
+ vq_path = hf_hub_download(
117
+ "fun-research/TiTok",
118
+ "maskgit-vqgan-imagenet-f16-256.bin",
119
+ revision="ab646ed225080a3acb7c78440a574d7f67f16fa7",
120
+ )
121
+ model = load_stellar(
122
+ "stellar-b16", purpose="reconstruct", vq_model=vq_path,
123
+ revision=revision, device=device,
124
+ )
125
+ with torch.no_grad():
126
+ features = model.encode(image)
127
+ reconstruction = model.reconstruct(features)
128
+ pixels = reconstruction["reconstruction"]
129
+ to_pil_image(pixels[0].cpu()).save("reconstruction.png")
130
+ ```
131
 
132
+ `reconstruct` takes **features, not an image**. You can also pass
133
+ `model.reconstruct(features["sparse"], features["spatial"])`. Both factors are
134
+ needed: semantic tokens alone do not specify spatial layout. B/L return RGB
135
+ 224x224 pixels; H returns 256x256 pixels. Outputs are in `[0, 1]`, with VQ token
136
+ IDs in `reconstruction["tokens"]` and logits in `reconstruction["logits"]`.
137
+ VQ decoding uses argmax and a frozen tokenizer, so this is not a differentiable
138
+ RGB decoder for end-to-end pixel losses.
139
 
140
  ### What the model returns
141
 
 
150
  `B` = batch, `K` = number of sparse tokens, `P` = number of patches (196 for /16 at
151
  224², 256 for /14), `D` = embedding dim (768 / 1024 / 1280 for B / L / H).
152
 
153
+ ## Training and evaluation
154
 
155
+ The GitHub repository provides image-folder training, checkpoint resume, and
156
+ classification, segmentation and reconstruction probes. Install the training
157
+ dependencies and prepare your data using the
158
+ [usage guide](https://github.com/microsoft/STELLAR/blob/main/docs/usage.md).
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
159
 
160
+ | Goal | Configuration or command |
161
+ | :--- | :--- |
162
+ | Pretrain on ImageNet or custom images | [configs/stellar.yaml](https://github.com/microsoft/STELLAR/blob/main/configs/stellar.yaml) |
163
+ | Train a classification probe | [configs/eval_cls.yaml](https://github.com/microsoft/STELLAR/blob/main/configs/eval_cls.yaml) |
164
+ | Train a segmentation probe | [configs/eval_seg.yaml](https://github.com/microsoft/STELLAR/blob/main/configs/eval_seg.yaml) |
165
+ | Train a reconstruction probe | [configs/eval_recon.yaml](https://github.com/microsoft/STELLAR/blob/main/configs/eval_recon.yaml) |
166
+ | Measure released-decoder rFID/LPIPS | [Reconstruction metrics](https://github.com/microsoft/STELLAR/blob/main/docs/usage.md#reconstruction-metrics) |
167
+
168
+ See the [configuration guide](https://github.com/microsoft/STELLAR#configuration-guide)
169
+ for data paths, batch size, devices, learning rate and model settings. Evaluation
170
+ recipes **train new heads on frozen features**; testing requires a trained head
171
+ checkpoint. The default pretraining recipe starts from MAE, not these Hub weights.
172
+
173
+ For a custom continued-pretraining loop, `load_stellar(..., purpose="pretrain",
174
+ vq_model=...)` loads all trained modules in train mode. Supply the multi-crop batch
175
+ contract described in the usage guide. This is **weights-only initialization**,
176
+ not an optimizer resume; use a full Lightning checkpoint with `scratch.resume`
177
+ to resume an interrupted training run.
178
 
179
  ## Model details
180
 
181
  - **Architecture:** ViT encoder (MAE-initialized) + learned sparse latent queries with
182
  spatial–semantic factorization.
183
  - **Pretraining data:** ImageNet-1K (self-supervised; labels not used).
184
+ - **Input:** RGB images in `[0, 1]` at 224x224. ImageNet mean/std
185
  normalization is applied **inside** the model — pass raw `[0, 1]` images.
186
  - **Weights:** the complete set of trained STELLAR modules (encoder, sparse tokens,
187
  projections, reconstruction decoder, and clustering heads), stored in `safetensors`.
 
191
 
192
  ## Intended uses & limitations
193
 
194
+ - **Supported research uses:** compact visual features for recognition,
195
+ segmentation, reconstruction, retrieval experiments and representation analysis.
196
+ - **Compression is not encoder acceleration:** the ViT still processes dense
197
+ patches. Fewer output latent values do not imply a 90% encoder or VLM speedup.
198
+ - **Tokens are learned concepts:** they are not guaranteed objects, class IDs,
199
+ temporal identities or pretrained language-aligned embeddings.
200
+ - **Domain shift:** ImageNet-pretrained features may require adaptation and
201
+ held-out evaluation on medical, satellite or other substantially different data.
202
+ - **Deployment:** these are research artifacts, not safety-tested models for
203
  production decision-making.
204
 
205
+ ### Connections to VLMs and latent generation
206
+
207
+ Sparse semantic tokens provide a compact input to a learned language-model
208
+ adapter. Appendix B.3 evaluates alignment to a CLIP text tower using a trained
209
+ probe; **no pretrained VLM adapter or captioning system is released**.
210
+
211
+ The paired semantic and spatial factors may also be studied as latents for
212
+ representation autoencoders (RAEs) or latent diffusion. **We have not evaluated
213
+ RAE-style generation and make no generation-quality claims.** Reconstruction
214
+ results do not establish unconditional or text-to-image generation performance.
215
+
216
  ## Citation
217
 
218
  ```bibtex
 
221
  author = {Zhao, Theodore Zhengde and Kiblawi, Sid and Yang, Jianwei and Usuyama, Naoto and Tan, Reuben and Codella, Noel C and Naumann, Tristan and Poon, Hoifung and Wei, Mu},
222
  booktitle = {International Conference on Machine Learning (ICML)},
223
  year = {2026},
224
+ url = {https://openreview.net/pdf?id=ysOOfySED6},
225
  }
226
  ```
227
 
228
  ## License
229
 
230
+ Released under the [MIT License](LICENSE).