File size: 1,383 Bytes
c6be157
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
"""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()