Spaces:
Running
Running
pixagram-search v4 embeddings (deploy.sh)
Browse files- README.md +180 -7
- app.py +153 -0
- handler.py +22 -0
- loadtest.py +77 -0
- requirements.txt +11 -0
- siglip.py +275 -0
README.md
CHANGED
|
@@ -1,13 +1,186 @@
|
|
| 1 |
---
|
| 2 |
-
title:
|
| 3 |
-
emoji:
|
| 4 |
-
colorFrom:
|
| 5 |
-
colorTo:
|
| 6 |
sdk: gradio
|
| 7 |
-
sdk_version: 6.
|
| 8 |
-
python_version: '3.13'
|
| 9 |
app_file: app.py
|
| 10 |
pinned: false
|
|
|
|
|
|
|
| 11 |
---
|
| 12 |
|
| 13 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
---
|
| 2 |
+
title: pixagram-search embeddings
|
| 3 |
+
emoji: 🟪
|
| 4 |
+
colorFrom: purple
|
| 5 |
+
colorTo: indigo
|
| 6 |
sdk: gradio
|
| 7 |
+
sdk_version: 6.28.0
|
|
|
|
| 8 |
app_file: app.py
|
| 9 |
pinned: false
|
| 10 |
+
license: apache-2.0
|
| 11 |
+
short_description: SigLIP image/text embeddings for the Pixagram search engine
|
| 12 |
---
|
| 13 |
|
| 14 |
+
# SigLIP embedding Space for pixagram-search
|
| 15 |
+
|
| 16 |
+
One model, two towers. The Cloudflare Worker sends artwork PNGs here while indexing and the
|
| 17 |
+
query text here at search time; both land in the same vector space, which is what makes
|
| 18 |
+
text → image search work. `app.py` is the Space entry point; it serves:
|
| 19 |
+
|
| 20 |
+
| route | |
|
| 21 |
+
|---|---|
|
| 22 |
+
| `POST /embed` | `{"inputs": {"images": ["<base64>", …], "texts": ["…", …]}}` → `{"model", "dim", "embeddings": [[…], …], "calibration", "max_num_patches"}` (images first, then texts; L2-normalised; `max_num_patches` for NaFlex models) |
|
| 23 |
+
| `GET /health` | `{"ok", "model", "dim", "ready", "stub", "calibration", "max_num_patches"}` — `ready` flips to true once the weights are loaded |
|
| 24 |
+
| `GET /` | Gradio page: embed a text or an image, score an image against captions |
|
| 25 |
+
|
| 26 |
+
`calibration` is the model's learned `{"logit_scale", "logit_bias"}`:
|
| 27 |
+
`sigmoid(logit_scale · cos + logit_bias)` is SigLIP's own text↔image match probability. The v3
|
| 28 |
+
Worker stores it (KV `calib:<model>`) and uses it until it has a background sample of the
|
| 29 |
+
corpus; older Spaces without the field still work (the Worker knows the SigLIP 2 base values).
|
| 30 |
+
|
| 31 |
+
## Deploy as a Space
|
| 32 |
+
|
| 33 |
+
1. Create a Space (SDK: **Gradio**, hardware: **CPU basic** is enough — a SigLIP-base pass on
|
| 34 |
+
one artwork takes ~0.25 s on 2 vCPU, ~0.5 s for NaFlex at 576 patches). The Space computes
|
| 35 |
+
one embedding per CPU at once and queues the rest; the Worker reads that number from
|
| 36 |
+
`/health` (`concurrency`) every ten minutes and sends as many per consumer invocation, so
|
| 37 |
+
bigger hardware is used with nothing else to change: on CPU upgrade (8 vCPU) it computes 8
|
| 38 |
+
at once instead of 2.
|
| 39 |
+
2. Upload this folder — `README.md` (the frontmatter above is the Space config), `app.py`,
|
| 40 |
+
`siglip.py`, `requirements.txt`; `handler.py` can stay, it is only used by Inference
|
| 41 |
+
Endpoints. From the repository root:
|
| 42 |
+
|
| 43 |
+
```bash
|
| 44 |
+
pip install -U huggingface_hub
|
| 45 |
+
hf auth login # older CLI: huggingface-cli login
|
| 46 |
+
hf upload <owner>/<space> hf . --repo-type space
|
| 47 |
+
```
|
| 48 |
+
|
| 49 |
+
`scripts/deploy.sh` does all of this for the v4 Space (`primerz/pixagram-search-v4`); the v3
|
| 50 |
+
repository did it for `primerz/pixagram-search-v3`.
|
| 51 |
+
3. Visibility:
|
| 52 |
+
* **Private** Space (recommended): every request must carry an HF token that can read the
|
| 53 |
+
Space — set that token as the Worker's `HF_TOKEN` secret. Leave `API_TOKEN` unset.
|
| 54 |
+
* **Public** Space: set a Space secret `API_TOKEN` to a long random string and use the same
|
| 55 |
+
string as the Worker's `HF_TOKEN`; `/embed` then rejects anything else with 401.
|
| 56 |
+
4. Worker: `npx wrangler secret put HF_TOKEN`, and set
|
| 57 |
+
`HF_EMBED_URL = "https://<owner>-<space-name>.hf.space/embed"` in `wrangler.jsonc`
|
| 58 |
+
(the host is `<owner>-<space-name>` with dots in the owner name replaced by dashes; the
|
| 59 |
+
Space settings page shows the exact "Direct URL").
|
| 60 |
+
5. Deploy the Worker, then `scripts/admin.sh reindex-all embed,text` to fill the vectors in.
|
| 61 |
+
|
| 62 |
+
Cold start: the SigLIP base checkpoints are ~1.5 GB; the Space answers `/health`
|
| 63 |
+
in seconds and `/embed` blocks until the weights are loaded (about 20 s once cached, a few
|
| 64 |
+
minutes on the first build). Free CPU Spaces **sleep after 48 h without traffic**; the Worker
|
| 65 |
+
treats the wake-up page as a transient error and retries with backoff, and search simply runs
|
| 66 |
+
without the semantic leg until the Space is back. If that gap matters, use paid CPU hardware
|
| 67 |
+
and set the sleep time to "never".
|
| 68 |
+
|
| 69 |
+
## Concurrency
|
| 70 |
+
|
| 71 |
+
The Space computes `MAX_CONCURRENCY` embeddings at once (default: one per CPU), each with
|
| 72 |
+
`CPUs / MAX_CONCURRENCY` PyTorch threads (`TORCH_THREADS` overrides), on threads that live as long
|
| 73 |
+
as the server. Search queries (texts only) have threads of their own and never wait behind
|
| 74 |
+
images being indexed. `/health` answers at once even under load.
|
| 75 |
+
|
| 76 |
+
Measured on 2 vCPU (Oct 2026): 24 NaFlex images at 576 patches from 4 clients, alone, then with
|
| 77 |
+
a search text every second:
|
| 78 |
+
|
| 79 |
+
| | images alone | images with searches | a search text meanwhile | `/health` meanwhile |
|
| 80 |
+
|---|---|---|---|---|
|
| 81 |
+
| before: one request at a time, 2 threads | 1.81 / s | 1.76 / s | median 1.25 s, worst 1.9 s | blocked (seconds) |
|
| 82 |
+
| **one per CPU, 1 thread each** | 1.89 / s | 1.77 / s | median 0.22 s, worst 0.29 s | 3 ms |
|
| 83 |
+
| one at a time, 2 threads | 1.55 / s | | | |
|
| 84 |
+
|
| 85 |
+
Two images at once with one thread each beat one image with two threads (1.89 against 1.55
|
| 86 |
+
images/s here), which is why the default is one per CPU. Fewer at once with more threads each
|
| 87 |
+
(`MAX_CONCURRENCY=1`) answers a single query sooner (~70 ms instead of ~125 ms for a text on 2
|
| 88 |
+
vCPU) at the cost of throughput.
|
| 89 |
+
|
| 90 |
+
`python3 hf/loadtest.py https://<owner>-<space>.hf.space <token>` measures a running Space the
|
| 91 |
+
same way (the token is the Worker's `HF_TOKEN`, `SPACE_API_TOKEN` in `~/.pixagram-search-v4.json`):
|
| 92 |
+
run it before and after changing the hardware.
|
| 93 |
+
|
| 94 |
+
## Troubleshooting
|
| 95 |
+
|
| 96 |
+
* Text embeddings taking 8-10 s on a CPU Space: `os.cpu_count()` inside the container reports
|
| 97 |
+
the host's cores (64-96) while the cgroup allows 2-8, and PyTorch spin-waits on the
|
| 98 |
+
oversubscribed thread pool. `siglip.py` sizes everything from the cgroup quota
|
| 99 |
+
(`effective_cpus()`); `/health` reports `cpus`, `concurrency` (embeddings at once) and
|
| 100 |
+
`threads` (PyTorch threads per embedding) so you can see what it picked.
|
| 101 |
+
|
| 102 |
+
* `[Errno 98] error while attempting to bind on address ('0.0.0.0', 7860)` in the build/run
|
| 103 |
+
log: Spaces set `GRADIO_SSR_MODE=True`, which makes Gradio start a Node SSR server on 7860
|
| 104 |
+
before uvicorn. `app.py` mounts Gradio with `ssr_mode=False` for that reason — keep it.
|
| 105 |
+
* `WARNING: Running pip as the 'root' user…` during the build is harmless; the Space image
|
| 106 |
+
installs requirements as root by design.
|
| 107 |
+
* `Your space is in error` on the direct URL: open the Space page → *Logs* (Build, then
|
| 108 |
+
Container); `/health` only answers once the container is running.
|
| 109 |
+
* Intermittent `502` HTML pages (fast, ~0.15 s, no `x-proxied-host` / `x-proxied-replica`
|
| 110 |
+
response header) while the Space shows *Running*: Hugging Face's edge fails before the
|
| 111 |
+
request reaches the container, so the Container log shows nothing. On 2026-09-28 this came in
|
| 112 |
+
windows of 10-50 s (up to 34 failures in a row). A quick retry does not bridge that. Search
|
| 113 |
+
falls back to full text, and the queue retries with backoff. Restart the Space. If it
|
| 114 |
+
persists, report it to HF with the `x-request-id` of a failed call, or move to paid hardware
|
| 115 |
+
or an Inference Endpoint (`handler.py`).
|
| 116 |
+
|
| 117 |
+
## Test
|
| 118 |
+
|
| 119 |
+
```bash
|
| 120 |
+
curl -s https://<owner>-<space>.hf.space/health
|
| 121 |
+
curl -s https://<owner>-<space>.hf.space/embed -H "Authorization: Bearer $HF_TOKEN" \
|
| 122 |
+
-H "Content-Type: application/json" -d '{"inputs":{"texts":["a swan on a lake at sunset"]}}' | jq '.dim'
|
| 123 |
+
```
|
| 124 |
+
|
| 125 |
+
Local run (any machine with Python 3.10+): `pip install -r requirements.txt && python app.py`,
|
| 126 |
+
then point the Worker's `.dev.vars` at `HF_EMBED_URL=http://127.0.0.1:7860/embed`.
|
| 127 |
+
`EMBED_STUB=1 python app.py` starts without weights and returns deterministic pseudo-vectors —
|
| 128 |
+
handy for wiring tests, never for production (`/health` reports `"stub": true`).
|
| 129 |
+
|
| 130 |
+
## Changing the model
|
| 131 |
+
|
| 132 |
+
The code default is `google/siglip-base-patch16-256-multilingual`, which is what the
|
| 133 |
+
production Space runs. Choose a model **per Space** with the `MODEL_ID` variable, so that
|
| 134 |
+
uploading `hf/` never changes production's model:
|
| 135 |
+
|
| 136 |
+
- the v2 Space (`primerz/pixagram-siglip2`) sets `MODEL_ID=google/siglip2-base-patch16-256`;
|
| 137 |
+
- the v3 Space (`primerz/pixagram-search-v3`) and the v4 Space (`primerz/pixagram-search-v4`,
|
| 138 |
+
created by this repository's `scripts/deploy.sh`) set `MODEL_ID=google/siglip2-base-patch16-naflex`
|
| 139 |
+
and `MAX_NUM_PATCHES=576`: v4 embeds exactly as v3 does.
|
| 140 |
+
|
| 141 |
+
### NaFlex
|
| 142 |
+
|
| 143 |
+
NaFlex checkpoints keep each image's aspect ratio. The processor resizes the image to the
|
| 144 |
+
largest size that fits `MAX_NUM_PATCHES` patches of 16×16 pixels (default 256), and batches mix
|
| 145 |
+
shapes through an attention mask. Every reply reports `max_num_patches`. The Worker refuses
|
| 146 |
+
image vectors whose budget differs from its `EMBED_PATCHES` (text vectors do not depend on it),
|
| 147 |
+
and re-embeds when `EMBED_PATCHES` changes. On the Pixagram corpus, 256 patches did worse than
|
| 148 |
+
the fixed 256 px model and 576 slightly better (README-V3.md, "SigLIP 2 NaFlex").
|
| 149 |
+
|
| 150 |
+
Any SigLIP / SigLIP 2 checkpoint works unchanged: fixed-resolution ones load as `SiglipModel`,
|
| 151 |
+
NaFlex ones as `Siglip2Model`. Stick to **multilingual** checkpoints, because queries arrive
|
| 152 |
+
in any language: every SigLIP 2 model, `google/siglip-base-patch16-256-multilingual` or
|
| 153 |
+
`google/siglip-so400m-patch16-256-i18n`. The other SigLIP 1 checkpoints have an English text
|
| 154 |
+
tower.
|
| 155 |
+
|
| 156 |
+
Measured on the 106 Pixagram artworks (+500 pixel-art distractors, 157 queries in EN/FR/DE/JA,
|
| 157 |
+
2 vCPU, Sept 2026):
|
| 158 |
+
|
| 159 |
+
| model | dim | ms / image | ms / query | notes |
|
| 160 |
+
|---|---|---|---|---|
|
| 161 |
+
| siglip-base-patch16-256-multilingual | 768 | 310 | 90 | production |
|
| 162 |
+
| **siglip2-base-patch16-256** | 768 | 315 | 87 | v2 Space; same cost; EN/FR on par, DE/JA better |
|
| 163 |
+
| siglip2-base-patch16-384 / -naflex (256 patches) | 768 | 650 / 305 | 85 | no gain over base-256 |
|
| 164 |
+
| **siglip2-base-patch16-naflex, 576 patches** | 768 | ~530 | 85 | v3 Space: small gain, measured in Oct 2026 on 133 artworks (README-V3.md) |
|
| 165 |
+
| siglip2-large-patch16-256 | 1024 | 980 | 290 | no gain |
|
| 166 |
+
| siglip2-so400m-patch16-256 | 1152 | 1270 | 390 | small gain (JA) at 4× the CPU |
|
| 167 |
+
| siglip-so400m-patch16-256-i18n | 1152 | 1280 | 395 | best measured, 4× the CPU |
|
| 168 |
+
|
| 169 |
+
Switching a stack in place, same dimension (768 → 768): keep the Vectorize index. Set the
|
| 170 |
+
Space's `MODEL_ID`, then set `EMBED_MODEL` to the same id in that stack's wrangler config
|
| 171 |
+
and deploy. The Worker refuses vectors whose `model` differs from `EMBED_MODEL`, and
|
| 172 |
+
re-embeds artworks whose `embed_model` differs (the 10-minute sweeper picks them up; to do it
|
| 173 |
+
at once: `scripts/admin.sh reindex-all embed,text`). Then refresh the background sample
|
| 174 |
+
(`scripts/admin.sh background`), since z-scores are per model.
|
| 175 |
+
|
| 176 |
+
A different dimension also needs **new Vectorize indexes** (with the metadata indexes from
|
| 177 |
+
`scripts/lib.sh`) that `VEC` and `VEC_TEXT` point at, and a new `EMBED_DIM`. Vectors from
|
| 178 |
+
different models must never be compared.
|
| 179 |
+
|
| 180 |
+
## Inference Endpoint instead of a Space
|
| 181 |
+
|
| 182 |
+
`handler.py` implements the `EndpointHandler` contract for a dedicated Inference Endpoint
|
| 183 |
+
(task *custom*, same `siglip.py`, same request/response shape). Use it when you want
|
| 184 |
+
autoscaling, private networking, or GPU without the Space UI. Then `HF_EMBED_URL` is the
|
| 185 |
+
endpoint URL itself and the Worker's `X-Scale-Up-Timeout: 600` header makes a scaled-to-zero
|
| 186 |
+
replica wake up within the request instead of returning 503.
|
app.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Hugging Face Space entry point (sdk: gradio): SigLIP image + text embeddings for pixagram-search.
|
| 3 |
+
|
| 4 |
+
Two faces on one server (port 7860):
|
| 5 |
+
POST /embed JSON API used by the Cloudflare Worker (see siglip.Embedder.handle for the shape)
|
| 6 |
+
GET /health {"ok", "model", "dim", "ready", "calibration"}
|
| 7 |
+
GET / Gradio demo page: embed a text or an image, compare an image with a caption
|
| 8 |
+
|
| 9 |
+
Space secrets / variables:
|
| 10 |
+
API_TOKEN optional shared secret; when set, /embed requires "Authorization: Bearer <API_TOKEN>"
|
| 11 |
+
(for a *private* Space leave it unset — HF already gates requests with your HF token)
|
| 12 |
+
MODEL_ID override the model (default google/siglip-base-patch16-256-multilingual, 768-d;
|
| 13 |
+
the v2 Space sets google/siglip2-base-patch16-256, the v3 Space
|
| 14 |
+
google/siglip2-base-patch16-naflex)
|
| 15 |
+
MAX_NUM_PATCHES NaFlex models: patches per image (default 256)
|
| 16 |
+
MAX_CONCURRENCY embeddings computed at once (default one per CPU; see siglip.py). /health
|
| 17 |
+
reports it as "concurrency", and the Worker's queue consumer sends that many
|
| 18 |
+
at once: a Space with more CPUs indexes faster without changing anything else.
|
| 19 |
+
"""
|
| 20 |
+
|
| 21 |
+
from __future__ import annotations
|
| 22 |
+
|
| 23 |
+
import asyncio
|
| 24 |
+
import os
|
| 25 |
+
import threading
|
| 26 |
+
import time
|
| 27 |
+
from concurrent.futures import ThreadPoolExecutor
|
| 28 |
+
from typing import Optional
|
| 29 |
+
|
| 30 |
+
import gradio as gr
|
| 31 |
+
import uvicorn
|
| 32 |
+
from fastapi import FastAPI, Header, HTTPException, Request
|
| 33 |
+
from fastapi.responses import JSONResponse
|
| 34 |
+
from PIL import Image
|
| 35 |
+
|
| 36 |
+
from siglip import MODEL_ID, Embedder, concurrency, effective_cpus, shared, threads_per_job, to_rgb
|
| 37 |
+
|
| 38 |
+
API_TOKEN = os.environ.get("API_TOKEN") or None
|
| 39 |
+
# Spaces expose port 7860 and set GRADIO_SERVER_PORT; PORT is honoured for other hosts.
|
| 40 |
+
PORT = int(os.environ.get("GRADIO_SERVER_PORT") or os.environ.get("PORT") or "7860")
|
| 41 |
+
HOST = os.environ.get("GRADIO_SERVER_NAME", "0.0.0.0")
|
| 42 |
+
|
| 43 |
+
embedder: Embedder = shared()
|
| 44 |
+
_started = time.time()
|
| 45 |
+
|
| 46 |
+
# CONCURRENCY embeddings at once, each on a thread of its own that lives as long as the server (a
|
| 47 |
+
# thread that runs PyTorch keeps its own OpenMP/MKL thread team: reusing the same threads keeps
|
| 48 |
+
# that to one team per slot). The event loop stays free for /health and for queueing the rest.
|
| 49 |
+
# Search queries (texts only) have threads of their own, so they never wait behind images being
|
| 50 |
+
# indexed; a text takes a fraction of an image's time.
|
| 51 |
+
CONCURRENCY = concurrency()
|
| 52 |
+
_images = ThreadPoolExecutor(max_workers=CONCURRENCY, thread_name_prefix="embed-image")
|
| 53 |
+
_texts = ThreadPoolExecutor(max_workers=CONCURRENCY, thread_name_prefix="embed-text")
|
| 54 |
+
|
| 55 |
+
|
| 56 |
+
def _pool(payload) -> ThreadPoolExecutor:
|
| 57 |
+
inputs = payload.get("inputs", payload) if isinstance(payload, dict) else payload
|
| 58 |
+
return _images if isinstance(inputs, dict) and inputs.get("images") else _texts
|
| 59 |
+
|
| 60 |
+
# Warm the model in the background so the first request does not pay the full load time.
|
| 61 |
+
threading.Thread(target=embedder.load, name="siglip-warmup", daemon=True).start()
|
| 62 |
+
|
| 63 |
+
api = FastAPI(title="pixagram-search embeddings", version="1.0")
|
| 64 |
+
|
| 65 |
+
|
| 66 |
+
@api.get("/health")
|
| 67 |
+
def health():
|
| 68 |
+
return {"ok": True, "model": embedder.model_id, "dim": embedder.dim, "ready": embedder.ready, "stub": embedder.stub,
|
| 69 |
+
"calibration": embedder.calibration, "max_num_patches": embedder.max_num_patches if embedder.naflex else None,
|
| 70 |
+
"concurrency": CONCURRENCY, "threads": threads_per_job(), "cpus": effective_cpus(),
|
| 71 |
+
"uptime_s": int(time.time() - _started)}
|
| 72 |
+
|
| 73 |
+
|
| 74 |
+
@api.post("/embed")
|
| 75 |
+
async def embed(request: Request, authorization: Optional[str] = Header(default=None)):
|
| 76 |
+
if API_TOKEN and authorization != f"Bearer {API_TOKEN}":
|
| 77 |
+
raise HTTPException(status_code=401, detail="unauthorized")
|
| 78 |
+
try:
|
| 79 |
+
payload = await request.json()
|
| 80 |
+
except Exception:
|
| 81 |
+
raise HTTPException(status_code=400, detail="body must be JSON")
|
| 82 |
+
try:
|
| 83 |
+
# embedder.handle blocks while the model finishes loading on a cold start; the Worker
|
| 84 |
+
# waits (its consumer has a 15-minute budget) rather than getting a 503.
|
| 85 |
+
return JSONResponse(await asyncio.get_running_loop().run_in_executor(_pool(payload), embedder.handle, payload))
|
| 86 |
+
except ValueError as e:
|
| 87 |
+
raise HTTPException(status_code=400, detail=str(e))
|
| 88 |
+
except Exception as e: # decoding errors etc.
|
| 89 |
+
raise HTTPException(status_code=422, detail=f"{type(e).__name__}: {e}")
|
| 90 |
+
|
| 91 |
+
|
| 92 |
+
# ---- Gradio demo -------------------------------------------------------------------------
|
| 93 |
+
|
| 94 |
+
|
| 95 |
+
def _vector_view(vec):
|
| 96 |
+
return {"dim": len(vec), "first_8": [round(v, 5) for v in vec[:8]], "norm": round(sum(v * v for v in vec) ** 0.5, 6)}
|
| 97 |
+
|
| 98 |
+
|
| 99 |
+
def ui_embed_text(text: str):
|
| 100 |
+
if not text or not text.strip():
|
| 101 |
+
return {"error": "type something"}
|
| 102 |
+
vec = embedder.embed_texts([text])[0]
|
| 103 |
+
return _vector_view(vec)
|
| 104 |
+
|
| 105 |
+
|
| 106 |
+
def ui_embed_image(img: Optional[Image.Image]):
|
| 107 |
+
if img is None:
|
| 108 |
+
return {"error": "upload an image"}
|
| 109 |
+
vec = embedder.embed_images([to_rgb(img)])[0]
|
| 110 |
+
return _vector_view(vec)
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
def ui_compare(img: Optional[Image.Image], texts: str):
|
| 114 |
+
if img is None or not texts.strip():
|
| 115 |
+
return {"error": "need an image and one caption per line"}
|
| 116 |
+
captions = [t.strip() for t in texts.splitlines() if t.strip()]
|
| 117 |
+
iv = embedder.embed_images([to_rgb(img)])[0]
|
| 118 |
+
tvs = embedder.embed_texts(captions)
|
| 119 |
+
scores = [(c, round(sum(a * b for a, b in zip(iv, tv)), 4)) for c, tv in zip(captions, tvs)]
|
| 120 |
+
scores.sort(key=lambda x: -x[1])
|
| 121 |
+
return {"cosine": scores}
|
| 122 |
+
|
| 123 |
+
|
| 124 |
+
with gr.Blocks(title="pixagram-search embeddings") as demo:
|
| 125 |
+
gr.Markdown(
|
| 126 |
+
f"## pixagram-search embeddings\n"
|
| 127 |
+
f"Model `{MODEL_ID}` — the same vectors the search engine stores in Vectorize. "
|
| 128 |
+
f"The Worker calls `POST /embed`; this page is for eyeballing the model on pixel art."
|
| 129 |
+
)
|
| 130 |
+
with gr.Tab("Text → vector"):
|
| 131 |
+
t_in = gr.Textbox(label="query text", placeholder="a swan on a lake at sunset")
|
| 132 |
+
t_btn = gr.Button("Embed")
|
| 133 |
+
t_out = gr.JSON(label="vector")
|
| 134 |
+
t_btn.click(ui_embed_text, inputs=t_in, outputs=t_out, api_name="embed_text")
|
| 135 |
+
with gr.Tab("Image → vector"):
|
| 136 |
+
i_in = gr.Image(type="pil", label="artwork (webp/png)")
|
| 137 |
+
i_btn = gr.Button("Embed")
|
| 138 |
+
i_out = gr.JSON(label="vector")
|
| 139 |
+
i_btn.click(ui_embed_image, inputs=i_in, outputs=i_out, api_name="embed_image")
|
| 140 |
+
with gr.Tab("Compare"):
|
| 141 |
+
c_img = gr.Image(type="pil", label="artwork")
|
| 142 |
+
c_txt = gr.Textbox(label="captions, one per line", lines=4, value="a swan on a lake at sunset\na space invader\na portrait of a woman\nabstract shapes")
|
| 143 |
+
c_btn = gr.Button("Score")
|
| 144 |
+
c_out = gr.JSON(label="cosine similarity per caption")
|
| 145 |
+
c_btn.click(ui_compare, inputs=[c_img, c_txt], outputs=c_out, api_name="compare")
|
| 146 |
+
|
| 147 |
+
# ssr_mode=False is essential on Spaces: the platform sets GRADIO_SSR_MODE=True, which would make
|
| 148 |
+
# mount_gradio_app start Gradio's Node SSR server on 7860 before uvicorn binds it
|
| 149 |
+
# ("[Errno 98] address already in use"). Client-side rendering needs no extra process.
|
| 150 |
+
app = gr.mount_gradio_app(api, demo, path="/", ssr_mode=False)
|
| 151 |
+
|
| 152 |
+
if __name__ == "__main__":
|
| 153 |
+
uvicorn.run(app, host=HOST, port=PORT)
|
handler.py
ADDED
|
@@ -0,0 +1,22 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
Alternative deployment: a Hugging Face *Inference Endpoint* (dedicated) instead of a Space.
|
| 3 |
+
Endpoints look for this file and the EndpointHandler class; the Space uses app.py.
|
| 4 |
+
Same request/response contract, so the Worker does not care which one it talks to.
|
| 5 |
+
"""
|
| 6 |
+
|
| 7 |
+
from __future__ import annotations
|
| 8 |
+
|
| 9 |
+
from typing import Any, Dict
|
| 10 |
+
|
| 11 |
+
from siglip import Embedder
|
| 12 |
+
|
| 13 |
+
|
| 14 |
+
class EndpointHandler:
|
| 15 |
+
def __init__(self, path: str = "") -> None:
|
| 16 |
+
self.embedder = Embedder().load()
|
| 17 |
+
|
| 18 |
+
def __call__(self, data: Dict[str, Any]) -> Dict[str, Any]:
|
| 19 |
+
try:
|
| 20 |
+
return self.embedder.handle(data)
|
| 21 |
+
except ValueError as e:
|
| 22 |
+
return {"error": str(e), "model": self.embedder.model_id, "dim": self.embedder.dim, "embeddings": []}
|
loadtest.py
ADDED
|
@@ -0,0 +1,77 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
"""Throughput of a running Space: 24 images from 4 clients at once (NaFlex-sized, 480x480), with a
|
| 3 |
+
search text every second, then the same without the texts. Run it before and after changing the
|
| 4 |
+
Space's hardware or MAX_CONCURRENCY.
|
| 5 |
+
|
| 6 |
+
python3 hf/loadtest.py https://primerz-pixagram-search-v4.hf.space [token]
|
| 7 |
+
|
| 8 |
+
The token is the Worker's HF_TOKEN (SPACE_API_TOKEN in ~/.pixagram-search-v4.json) when the
|
| 9 |
+
Space has an API_TOKEN; it is read from $SPACE_API_TOKEN when not given.
|
| 10 |
+
"""
|
| 11 |
+
import base64
|
| 12 |
+
import io
|
| 13 |
+
import json
|
| 14 |
+
import os
|
| 15 |
+
import sys
|
| 16 |
+
import threading
|
| 17 |
+
import time
|
| 18 |
+
import urllib.request
|
| 19 |
+
|
| 20 |
+
from PIL import Image
|
| 21 |
+
|
| 22 |
+
BASE = sys.argv[1].rstrip("/")
|
| 23 |
+
TOKEN = sys.argv[2] if len(sys.argv) > 2 else os.environ.get("SPACE_API_TOKEN", "")
|
| 24 |
+
buf = io.BytesIO()
|
| 25 |
+
Image.effect_noise((480, 480), 64).convert("RGB").save(buf, "PNG")
|
| 26 |
+
IMG = base64.b64encode(buf.getvalue()).decode()
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def post(body):
|
| 30 |
+
headers = {"content-type": "application/json", **({"authorization": f"Bearer {TOKEN}"} if TOKEN else {})}
|
| 31 |
+
req = urllib.request.Request(BASE + "/embed", data=json.dumps(body).encode(), headers=headers)
|
| 32 |
+
with urllib.request.urlopen(req, timeout=600) as r:
|
| 33 |
+
return json.load(r)
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
def run(with_texts: bool) -> str:
|
| 37 |
+
done, text_ms = False, []
|
| 38 |
+
|
| 39 |
+
def client():
|
| 40 |
+
for _ in range(6):
|
| 41 |
+
post({"inputs": {"images": [IMG]}})
|
| 42 |
+
|
| 43 |
+
def texts():
|
| 44 |
+
while not done:
|
| 45 |
+
time.sleep(1)
|
| 46 |
+
a = time.time()
|
| 47 |
+
post({"inputs": {"texts": ["a red dragon over a castle"]}})
|
| 48 |
+
text_ms.append((time.time() - a) * 1000)
|
| 49 |
+
|
| 50 |
+
t0 = time.time()
|
| 51 |
+
clients = [threading.Thread(target=client) for _ in range(4)]
|
| 52 |
+
tx = threading.Thread(target=texts) if with_texts else None
|
| 53 |
+
for c in clients:
|
| 54 |
+
c.start()
|
| 55 |
+
if tx:
|
| 56 |
+
tx.start()
|
| 57 |
+
for c in clients:
|
| 58 |
+
c.join()
|
| 59 |
+
wall = time.time() - t0
|
| 60 |
+
done = True
|
| 61 |
+
if tx:
|
| 62 |
+
tx.join()
|
| 63 |
+
out = f"{24 / wall:.2f} images/s"
|
| 64 |
+
if text_ms:
|
| 65 |
+
text_ms.sort()
|
| 66 |
+
out += f"; a search text meanwhile: median {text_ms[len(text_ms) // 2]:.0f} ms, worst {text_ms[-1]:.0f} ms"
|
| 67 |
+
return out
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
if __name__ == "__main__":
|
| 71 |
+
with urllib.request.urlopen(BASE + "/health", timeout=60) as r:
|
| 72 |
+
h = json.load(r)
|
| 73 |
+
print("health:", {k: h.get(k) for k in ("model", "cpus", "concurrency", "threads")})
|
| 74 |
+
post({"inputs": {"texts": ["warm up"]}})
|
| 75 |
+
post({"inputs": {"images": [IMG]}})
|
| 76 |
+
print("images alone: ", run(False))
|
| 77 |
+
print("images and searches:", run(True))
|
requirements.txt
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
torch>=2.2
|
| 2 |
+
# SigLIP 2 needs >= 4.49 (Gemma tokenizer config); tested with 5.17
|
| 3 |
+
transformers>=4.49
|
| 4 |
+
pillow>=10
|
| 5 |
+
# sentencepiece: tokenizer of the previous model (siglip-base-patch16-256-multilingual) — keep for rollback
|
| 6 |
+
sentencepiece>=0.2
|
| 7 |
+
protobuf>=4
|
| 8 |
+
# app.py (Space) — the gradio SDK image already ships gradio/fastapi/uvicorn; pins keep local runs identical
|
| 9 |
+
gradio>=5
|
| 10 |
+
fastapi>=0.110
|
| 11 |
+
uvicorn>=0.29
|
siglip.py
ADDED
|
@@ -0,0 +1,275 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""
|
| 2 |
+
SigLIP embedder shared by app.py (Space) and handler.py (Inference Endpoint).
|
| 3 |
+
|
| 4 |
+
One model, two towers: images for the indexing pipeline, texts for search queries — both land
|
| 5 |
+
in the same vector space, which is what makes text → image search work. Vectors are
|
| 6 |
+
L2-normalised so cosine similarity is a dot product (Vectorize metric: cosine).
|
| 7 |
+
|
| 8 |
+
Text is lowercased here, before tokenisation. SigLIP 2 was trained on lowercased text and its
|
| 9 |
+
Gemma tokenizer is case-sensitive: on Pixagram artworks, Title Case queries drop from 97 % to
|
| 10 |
+
3 % recall@1. The original multilingual SigLIP tokenizer lowercases on its own, so this is a
|
| 11 |
+
no-op for it. Doing it here covers every caller (search queries, blog résumés, the demo page).
|
| 12 |
+
|
| 13 |
+
NaFlex models (google/siglip2-*-naflex) keep each image's aspect ratio: the processor resizes it to
|
| 14 |
+
the largest size that fits MAX_NUM_PATCHES patches of 16x16 px, instead of squashing it to a fixed
|
| 15 |
+
square. A 274x183 artwork is seen as about 19x13 patches rather than distorted to 16x16.
|
| 16 |
+
|
| 17 |
+
Environment:
|
| 18 |
+
MODEL_ID Hub id of a SigLIP/CLIP-family model (default: multilingual SigLIP base, 768-d;
|
| 19 |
+
the v2 Space primerz/pixagram-siglip2 sets google/siglip2-base-patch16-256, the
|
| 20 |
+
v3 Space google/siglip2-base-patch16-naflex)
|
| 21 |
+
MAX_NUM_PATCHES NaFlex only: patch budget per image (default 256, the model's training default).
|
| 22 |
+
Reported in every reply so the Worker can check it matches EMBED_PATCHES.
|
| 23 |
+
MAX_BATCH images/texts per forward pass (default 32)
|
| 24 |
+
MAX_CONCURRENCY embeddings computed at once (default: one per CPU, effective_cpus()); each
|
| 25 |
+
gets CPUs / MAX_CONCURRENCY PyTorch threads. One per CPU gives the most
|
| 26 |
+
throughput (a small forward pass does not spread well over many threads);
|
| 27 |
+
fewer, with more threads each, answer a single query sooner.
|
| 28 |
+
TORCH_THREADS PyTorch threads per embedding, overriding CPUs / MAX_CONCURRENCY
|
| 29 |
+
EMBED_STUB "1" returns deterministic pseudo-vectors without loading any model — wiring tests only
|
| 30 |
+
"""
|
| 31 |
+
|
| 32 |
+
from __future__ import annotations
|
| 33 |
+
|
| 34 |
+
import base64
|
| 35 |
+
import hashlib
|
| 36 |
+
import io
|
| 37 |
+
import math
|
| 38 |
+
import os
|
| 39 |
+
import threading
|
| 40 |
+
from typing import Any, Dict, List, Optional
|
| 41 |
+
|
| 42 |
+
from PIL import Image
|
| 43 |
+
|
| 44 |
+
# Default = what the production Space runs. Change models per Space with the MODEL_ID variable,
|
| 45 |
+
# so pushing hf/ never switches production's model by accident.
|
| 46 |
+
MODEL_ID = os.environ.get("MODEL_ID", "google/siglip-base-patch16-256-multilingual")
|
| 47 |
+
MAX_BATCH = int(os.environ.get("MAX_BATCH", "32"))
|
| 48 |
+
MAX_NUM_PATCHES = int(os.environ.get("MAX_NUM_PATCHES", "256"))
|
| 49 |
+
MAX_TEXT_TOKENS = 64 # SigLIP text-tower context
|
| 50 |
+
STUB = os.environ.get("EMBED_STUB", "") == "1"
|
| 51 |
+
STUB_DIM = int(os.environ.get("EMBED_STUB_DIM", "768"))
|
| 52 |
+
|
| 53 |
+
|
| 54 |
+
def decode_image(b64: str) -> Image.Image:
|
| 55 |
+
"""base64 (optionally a data URI) → RGB PIL image, transparency composited over white."""
|
| 56 |
+
if b64.startswith("data:"):
|
| 57 |
+
b64 = b64.split(",", 1)[1]
|
| 58 |
+
raw = base64.b64decode(b64)
|
| 59 |
+
img = Image.open(io.BytesIO(raw))
|
| 60 |
+
img.load()
|
| 61 |
+
return to_rgb(img)
|
| 62 |
+
|
| 63 |
+
|
| 64 |
+
def to_rgb(img: Image.Image) -> Image.Image:
|
| 65 |
+
if img.mode in ("RGBA", "LA", "P"):
|
| 66 |
+
rgba = img.convert("RGBA")
|
| 67 |
+
bg = Image.new("RGBA", rgba.size, (255, 255, 255, 255))
|
| 68 |
+
img = Image.alpha_composite(bg, rgba)
|
| 69 |
+
return img.convert("RGB")
|
| 70 |
+
|
| 71 |
+
|
| 72 |
+
class Embedder:
|
| 73 |
+
"""Lazy-loading, thread-safe wrapper around a SigLIP model."""
|
| 74 |
+
|
| 75 |
+
def __init__(self, model_id: str = MODEL_ID, max_num_patches: int = MAX_NUM_PATCHES) -> None:
|
| 76 |
+
self.model_id = model_id
|
| 77 |
+
self._lock = threading.Lock()
|
| 78 |
+
self._model = None
|
| 79 |
+
self._processor = None
|
| 80 |
+
self.device = "cpu"
|
| 81 |
+
self.dim = STUB_DIM if STUB else 0
|
| 82 |
+
self.stub = STUB
|
| 83 |
+
# NaFlex: set on load when the image processor takes a patch budget
|
| 84 |
+
self.naflex = False
|
| 85 |
+
self.max_num_patches = max_num_patches
|
| 86 |
+
# sigmoid(logit_scale * cos + logit_bias) is the model's own text/image match probability
|
| 87 |
+
self.calibration: Optional[Dict[str, float]] = None
|
| 88 |
+
|
| 89 |
+
# ---- loading --------------------------------------------------------------------------
|
| 90 |
+
|
| 91 |
+
def load(self) -> "Embedder":
|
| 92 |
+
if self.stub or self._model is not None:
|
| 93 |
+
return self
|
| 94 |
+
with self._lock:
|
| 95 |
+
if self._model is not None:
|
| 96 |
+
return self
|
| 97 |
+
import torch
|
| 98 |
+
from transformers import AutoModel, AutoProcessor
|
| 99 |
+
|
| 100 |
+
self.device = "cuda" if torch.cuda.is_available() else "cpu"
|
| 101 |
+
# os.cpu_count() reports the HOST's cores inside a container (often 64-96 on Spaces)
|
| 102 |
+
# while the cgroup allows 2-8; oversubscribing PyTorch's thread pool that badly turns
|
| 103 |
+
# an 80 ms text embedding into 8-12 s of spin-waiting. Size the pool to what the
|
| 104 |
+
# container can actually run, shared by the embeddings computed at once.
|
| 105 |
+
torch.set_num_threads(threads_per_job())
|
| 106 |
+
model = AutoModel.from_pretrained(self.model_id).to(self.device).eval()
|
| 107 |
+
self._processor = AutoProcessor.from_pretrained(self.model_id)
|
| 108 |
+
self.naflex = hasattr(getattr(self._processor, "image_processor", None), "max_num_patches")
|
| 109 |
+
cfg = model.config
|
| 110 |
+
self.dim = int(getattr(cfg, "projection_dim", 0) or getattr(cfg.text_config, "projection_size", 0) or cfg.text_config.hidden_size)
|
| 111 |
+
# SigLIP's learned temperature and bias (exp() of the stored log-scale). CLIP-family models
|
| 112 |
+
# without a bias report none.
|
| 113 |
+
try:
|
| 114 |
+
scale = float(model.logit_scale.exp().item())
|
| 115 |
+
bias = float(model.logit_bias.item()) if getattr(model, "logit_bias", None) is not None else None
|
| 116 |
+
self.calibration = {"logit_scale": round(scale, 4), "logit_bias": round(bias, 4)} if bias is not None else None
|
| 117 |
+
except Exception:
|
| 118 |
+
self.calibration = None
|
| 119 |
+
self._model = model
|
| 120 |
+
# First inference pays one-off kernel/allocator warm-up (~2 s); do it here, not on a user query.
|
| 121 |
+
try:
|
| 122 |
+
self.embed_texts(["warm up"])
|
| 123 |
+
self.embed_images([Image.new("RGB", (32, 32), (128, 128, 128))])
|
| 124 |
+
except Exception:
|
| 125 |
+
pass
|
| 126 |
+
return self
|
| 127 |
+
|
| 128 |
+
@property
|
| 129 |
+
def ready(self) -> bool:
|
| 130 |
+
return self.stub or self._model is not None
|
| 131 |
+
|
| 132 |
+
# ---- embedding ------------------------------------------------------------------------
|
| 133 |
+
|
| 134 |
+
def embed_images(self, images: List[Image.Image]) -> List[List[float]]:
|
| 135 |
+
if self.stub:
|
| 136 |
+
return [_stub_vector(hashlib.sha256(im.tobytes()).hexdigest()) for im in images]
|
| 137 |
+
self.load()
|
| 138 |
+
import torch
|
| 139 |
+
|
| 140 |
+
out: List[List[float]] = []
|
| 141 |
+
# NaFlex: images of any aspect ratio share a batch; the processor pads them to the patch
|
| 142 |
+
# budget and passes the attention mask and each image's patch grid to the model.
|
| 143 |
+
kw = {"max_num_patches": self.max_num_patches} if self.naflex else {}
|
| 144 |
+
with torch.inference_mode():
|
| 145 |
+
for i in range(0, len(images), MAX_BATCH):
|
| 146 |
+
batch = [to_rgb(im) for im in images[i : i + MAX_BATCH]]
|
| 147 |
+
inputs = self._processor(images=batch, return_tensors="pt", **kw).to(self.device)
|
| 148 |
+
feats = self._model.get_image_features(**inputs)
|
| 149 |
+
feats = _pooled(feats)
|
| 150 |
+
feats = torch.nn.functional.normalize(feats, dim=-1)
|
| 151 |
+
out.extend(feats.cpu().float().tolist())
|
| 152 |
+
return out
|
| 153 |
+
|
| 154 |
+
def embed_texts(self, texts: List[str]) -> List[List[float]]:
|
| 155 |
+
if self.stub:
|
| 156 |
+
return [_stub_vector(hashlib.sha256(t.strip().lower().encode()).hexdigest()) for t in texts]
|
| 157 |
+
self.load()
|
| 158 |
+
import torch
|
| 159 |
+
|
| 160 |
+
out: List[List[float]] = []
|
| 161 |
+
with torch.inference_mode():
|
| 162 |
+
for i in range(0, len(texts), MAX_BATCH):
|
| 163 |
+
# Lowercase: SigLIP 2 was trained on lowercased text (see module docstring).
|
| 164 |
+
batch = [t.lower() if t.strip() else " " for t in texts[i : i + MAX_BATCH]]
|
| 165 |
+
# SigLIP was trained with padding="max_length"; keep it so query vectors match.
|
| 166 |
+
inputs = self._processor(
|
| 167 |
+
text=batch, padding="max_length", truncation=True, max_length=MAX_TEXT_TOKENS, return_tensors="pt"
|
| 168 |
+
).to(self.device)
|
| 169 |
+
feats = self._model.get_text_features(**inputs)
|
| 170 |
+
feats = _pooled(feats)
|
| 171 |
+
feats = torch.nn.functional.normalize(feats, dim=-1)
|
| 172 |
+
out.extend(feats.cpu().float().tolist())
|
| 173 |
+
return out
|
| 174 |
+
|
| 175 |
+
# ---- request handling shared by the Space route and the Endpoint handler -----------------
|
| 176 |
+
|
| 177 |
+
def handle(self, payload: Dict[str, Any]) -> Dict[str, Any]:
|
| 178 |
+
"""
|
| 179 |
+
{"inputs": {"images": [b64, ...], "texts": [str, ...]}} (either key optional)
|
| 180 |
+
→ {"model", "dim", "embeddings": [image vectors..., text vectors...], "calibration": {"logit_scale", "logit_bias"}}
|
| 181 |
+
"""
|
| 182 |
+
inputs = payload.get("inputs", payload) if isinstance(payload, dict) else payload
|
| 183 |
+
if isinstance(inputs, str):
|
| 184 |
+
inputs = {"texts": [inputs]}
|
| 185 |
+
if not isinstance(inputs, dict):
|
| 186 |
+
raise ValueError("body must be {\"inputs\": {\"images\": [...], \"texts\": [...]}}")
|
| 187 |
+
images_b64 = inputs.get("images") or []
|
| 188 |
+
texts = inputs.get("texts") or []
|
| 189 |
+
if isinstance(images_b64, str):
|
| 190 |
+
images_b64 = [images_b64]
|
| 191 |
+
if isinstance(texts, str):
|
| 192 |
+
texts = [texts]
|
| 193 |
+
if not images_b64 and not texts:
|
| 194 |
+
raise ValueError("provide inputs.images (base64) and/or inputs.texts")
|
| 195 |
+
if len(images_b64) + len(texts) > 256:
|
| 196 |
+
raise ValueError("at most 256 items per request")
|
| 197 |
+
|
| 198 |
+
embeddings: List[List[float]] = []
|
| 199 |
+
if images_b64:
|
| 200 |
+
embeddings.extend(self.embed_images([decode_image(b) for b in images_b64]))
|
| 201 |
+
if texts:
|
| 202 |
+
embeddings.extend(self.embed_texts([str(t) for t in texts]))
|
| 203 |
+
out: Dict[str, Any] = {"model": self.model_id if not self.stub else "stub", "dim": len(embeddings[0]) if embeddings else self.dim, "embeddings": embeddings}
|
| 204 |
+
if self.calibration:
|
| 205 |
+
out["calibration"] = self.calibration
|
| 206 |
+
if self.naflex:
|
| 207 |
+
out["max_num_patches"] = self.max_num_patches
|
| 208 |
+
return out
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
def effective_cpus(cap: int = 16) -> int:
|
| 212 |
+
"""CPUs this process may really use: affinity mask ∩ cgroup quota, capped."""
|
| 213 |
+
n = os.cpu_count() or 1
|
| 214 |
+
try:
|
| 215 |
+
n = min(n, len(os.sched_getaffinity(0)))
|
| 216 |
+
except Exception:
|
| 217 |
+
pass
|
| 218 |
+
try: # cgroup v2
|
| 219 |
+
quota, period = open("/sys/fs/cgroup/cpu.max").read().split()[:2]
|
| 220 |
+
if quota != "max":
|
| 221 |
+
n = min(n, max(1, int(int(quota) / int(period))))
|
| 222 |
+
except Exception:
|
| 223 |
+
try: # cgroup v1
|
| 224 |
+
quota = int(open("/sys/fs/cgroup/cpu/cpu.cfs_quota_us").read())
|
| 225 |
+
period = int(open("/sys/fs/cgroup/cpu/cpu.cfs_period_us").read())
|
| 226 |
+
if quota > 0:
|
| 227 |
+
n = min(n, max(1, quota // period))
|
| 228 |
+
except Exception:
|
| 229 |
+
pass
|
| 230 |
+
return max(1, min(n, cap))
|
| 231 |
+
|
| 232 |
+
|
| 233 |
+
def concurrency() -> int:
|
| 234 |
+
"""Embeddings computed at once: MAX_CONCURRENCY, else one per CPU."""
|
| 235 |
+
env = os.environ.get("MAX_CONCURRENCY", "")
|
| 236 |
+
return max(1, int(env)) if env.isdigit() else effective_cpus()
|
| 237 |
+
|
| 238 |
+
|
| 239 |
+
def threads_per_job() -> int:
|
| 240 |
+
"""PyTorch threads per embedding: TORCH_THREADS, else the CPUs shared by the embeddings at once."""
|
| 241 |
+
env = os.environ.get("TORCH_THREADS", "")
|
| 242 |
+
if env.isdigit() and int(env) > 0:
|
| 243 |
+
return int(env)
|
| 244 |
+
return max(1, effective_cpus() // concurrency())
|
| 245 |
+
|
| 246 |
+
|
| 247 |
+
def _pooled(feats):
|
| 248 |
+
"""transformers ≥5 may return a ModelOutput; older versions return the tensor directly."""
|
| 249 |
+
if hasattr(feats, "pooler_output") and feats.pooler_output is not None:
|
| 250 |
+
return feats.pooler_output
|
| 251 |
+
if hasattr(feats, "last_hidden_state") and not hasattr(feats, "shape"):
|
| 252 |
+
return feats.last_hidden_state
|
| 253 |
+
return feats
|
| 254 |
+
|
| 255 |
+
|
| 256 |
+
def _stub_vector(seed_hex: str) -> List[float]:
|
| 257 |
+
"""Deterministic unit vector from a hash — lets the Worker ↔ Space wiring be tested without weights."""
|
| 258 |
+
vals: List[float] = []
|
| 259 |
+
h = seed_hex
|
| 260 |
+
while len(vals) < STUB_DIM:
|
| 261 |
+
h = hashlib.sha256(h.encode()).hexdigest()
|
| 262 |
+
vals.extend(int(h[i : i + 2], 16) / 127.5 - 1.0 for i in range(0, 64, 2))
|
| 263 |
+
vals = vals[:STUB_DIM]
|
| 264 |
+
norm = math.sqrt(sum(v * v for v in vals)) or 1.0
|
| 265 |
+
return [v / norm for v in vals]
|
| 266 |
+
|
| 267 |
+
|
| 268 |
+
_shared: Optional[Embedder] = None
|
| 269 |
+
|
| 270 |
+
|
| 271 |
+
def shared() -> Embedder:
|
| 272 |
+
global _shared
|
| 273 |
+
if _shared is None:
|
| 274 |
+
_shared = Embedder()
|
| 275 |
+
return _shared
|