needle3-gpu / README.md
dkappe's picture
Upload README.md with huggingface_hub
2acf1e8 verified
|
Raw History Blame Contribute Delete
4.11 kB
---
title: needle3 logits
emoji: 🌡
colorFrom: green
colorTo: gray
sdk: gradio
sdk_version: 5.49.1
app_file: app.py
pinned: false
license: apache-2.0
short_description: Jev-shaped endpoint on Needle 3 logits, ZeroGPU-capable
---
# needle3 logits (Gradio SDK + ZeroGPU)
A Jev-shaped endpoint backed by [Cactus Needle 3](https://huggingface.co/Cactus-Compute/needle3),
reading its **JAX path** so real next-token logits are available. Runs on **ZeroGPU**.
## Verified working
- **ZeroGPU allocated** β€” hardware `zero-a10g`, Space `RUNNING`.
- **JAX on ZeroGPU works** β€” the `GPU probe` tab reports `backend: gpu`,
`devices: ["cuda:0"]`, and completes a real computation. This is *not* a documented-supported
combination (ZeroGPU is PyTorch-shaped), so it was verified rather than assumed.
- **Real distributions** β€” `choice` returns genuine probabilities, e.g. a billing ticket:
`billing 0.877 / technical 0.004 / account 0.120`.
## Honest status
**The plumbing works. Accuracy does not yet.**
- Second test case ("I cannot log in, my password is rejected") returns
`billing 0.709 / account 0.285` β€” **wrong**, should be `account`. So `choice` is roughly
2/3 on unambiguous cases, not production quality.
- `noul` (yes/no) did not discriminate on matched positive/negative pairs.
- Probabilities are **uncalibrated** β€” nothing shows a stated 0.8 is right 80% of the time.
Do not build a threshold on them.
## Why the JAX path
Needle 3 ships two runtimes, and only one yields probabilities:
| Path | Exposes | Logits? |
|---|---|---|
| Native C engine (`libneedle.so` + `needle3.cact`) | `needle_init/complete/embed/load/reset` | ❌ no |
| **JAX** (`needle/model/run.py`, `needle3.safetensors`) | `logits = decode_fn(params, buffer)[0, pos]` | βœ… **yes** |
The native engine's `confidence` head returned a constant `1.0` in local testing.
## ZeroGPU implementation notes
Four things were required, each of which failed first:
1. **A NEW Space.** ZeroGPU requires the Gradio SDK, and a Space's SDK is immutable β€” a Docker
Space cannot be converted. This replaced an earlier Docker build (`dkappe/needle3-logits`),
which was deleted once this Gradio version was verified. Its source is still in the repo
beside this one as reference, but nothing runs it.
2. **Hardware request via REST.** `hf spaces settings <id> --hardware zero-a10g` errors on this
CLI version; a `PUT /api/spaces/<id>/settings` with `{"hardware":"zero-a10g"}` works.
3. **`config.dtype = "float32"` before building the network.** The checkpoint declares
`dtype=bfloat16`, and every layer is constructed with it. On this GPU that aborts with
`Unsupported conversion from bf16 to f16` / `LLVM ERROR: Unsupported rounding mode for
conversion.` Casting the *parameters* alone is not enough β€” the model's own dtype must
change too. `NEEDLE_CAST=bfloat16` reverts this.
4. **Checkpoint fetch at module scope, inside `if __name__ == "__main__"`.** Gradio executes
`app_file` as `__main__`, so a bootstrap in an `else:` branch never runs; and network access
inside a `@spaces.GPU` function is restricted.
Also: `XLA_PYTHON_CLIENT_PREALLOCATE=false` (JAX otherwise claims ~75% of a shared GPU slice),
and the health check deliberately does **not** import JAX β€” outside a `@spaces.GPU` function
there is no GPU attached, so it would fail with `No visible GPU devices`.
## Endpoints
| Endpoint | What it does |
|---|---|
| `/choice` | categorical pick with a full distribution |
| `/noul` | P(yes) as one number |
| `/score` | ordered rating; can land between levels |
| `/methods` | scoring-method comparison |
| `/selftest` | matched positive/negative probes |
| `/gpu_probe` | device list + a real computation (GPU proof) |
| `/health` | files/env only, no JAX |
## Caveats
- Experimental; not affiliated with TypeSafe AI or Cactus Compute.
- Free ZeroGPU quota is small (~5 min/day for a free account), so heavy experimentation will
exhaust it.
- Accuracy figures come from three hand-chosen cases β€” a smoke test, not a benchmark.
- Apache-2.0, matching the model.