"""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()