Download examples/classify.py from zeechimp/gnomon: direct link, hf CLI and curl.
- Browser
- Download file 1.38 kB
-
https://huggingface.co/zeechimp/gnomon/resolve/main/examples/classify.py
- Command line
-
hf download hf://zeechimp/gnomon/examples/classify.py
-
curl -L -o classify.py https://huggingface.co/zeechimp/gnomon/resolve/main/examples/classify.py
1.38 kB
| """Compile a small classifier into an LDC and run at inference speed.""" | |
| import os | |
| import sys | |
| import time | |
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) | |
| import numpy as np | |
| import jax | |
| import jax.numpy as jnp | |
| from gnomon import compile_ldc | |
| def main(): | |
| rng = np.random.default_rng(0) | |
| N, D, C = 20_000, 32, 3 | |
| X = rng.standard_normal((N, D)).astype(np.float32) | |
| W_true = rng.standard_normal((D, C)).astype(np.float32) | |
| labels = (X @ W_true).argmax(axis=1).astype(np.int32) | |
| X_train = X[:10_000] | |
| Y_train = np.eye(C, dtype=np.float32)[labels[:10_000]] | |
| X_test = X[10_000:] | |
| y_test = labels[10_000:] | |
| print("Compiling LDC...") | |
| t0 = time.perf_counter() | |
| ldc = compile_ldc(X_train, Y_train, k=4, max_leaves=64, projection_steps=300) | |
| print(f" compiled in {time.perf_counter() - t0:.2f}s, {ldc.n_leaves} leaves") | |
| # Warmup | |
| _ = ldc.predict(jnp.asarray(X_test[:8])) | |
| jax.block_until_ready(_) | |
| t0 = time.perf_counter() | |
| preds = ldc.predict(jnp.asarray(X_test)) | |
| jax.block_until_ready(preds) | |
| dt = time.perf_counter() - t0 | |
| pred_labels = np.asarray(preds.argmax(-1)) | |
| acc = (pred_labels == y_test).mean() | |
| print(f"Test accuracy: {acc:.4f}") | |
| print(f"Throughput: {len(X_test) / dt:,.0f} queries/s") | |
| if __name__ == "__main__": | |
| main() |