gnomon / examples /classify.py
zeechimp's picture
Upload 15 files
c6be157 verified
Raw History Blame Contribute Delete
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()