Backends

cvi comes with some optimizations in the form of various backends that you can switch between for faster performance depending on your use-case.

Optional optimization coverage

Index

Numba batch

Numba sample updates

Numba remove/merge

JAX batch

JAX streaming

CH

Yes

Unchanged

Unchanged

Yes

Capacity

CONN

Unavailable

Unavailable

Unavailable

Unavailable

Unavailable

cSIL

Yes

Unchanged

Unchanged

Unavailable

Unavailable

DB

Yes

Distances

Distances

Unavailable

Unavailable

GD43

Yes

Distances

Distances

Unavailable

Unavailable

GD53

Yes

Unchanged

Unchanged

Unavailable

Unavailable

PS

Yes

Distances

Distances

Unavailable

Unavailable

rCIP

Unavailable

Unavailable

Unavailable

Unavailable

Unavailable

WB

Yes

Unchanged

Unchanged

Yes

Capacity

XB

Yes

Distances

Distances

Yes

Capacity

Yes means compiled batch kernels are available; it does not mean that every operation is compiled. Distances means centroid-distance kernels are compiled while other update logic remains in Python/NumPy. Unchanged means the operation is supported with backend="numba" but uses its existing NumPy implementation. Unavailable means the index rejects that numerical backend.

Capacity means JAX sample additions and update_many chunks require a positive capacity at construction. CH, WB, and XB also provide functional JAX batch and streaming APIs. All JAX modes require x64, and none support remove or merge. Numba retains remove/merge support for all eight supported indices, including the operations marked Unchanged.

NumPy remains the default for every index. CONN’s model_type selects its prototype learner separately from numerical backend selection; it does not enable the optional Numba or JAX numerical backends.

Backend availability does not guarantee a speedup. The sections below describe which kernels are compiled and when compilation or transfer costs matter.

Optional CPU acceleration

Install the Numba extra to enable compiled numerical kernels:

python -m pip install "cvi[numba]"

For a development checkout, use python -m pip install -e ".[numba]". Select the backend when constructing an index:

index = cvi.XB(backend="numba")
value = index.get_cvi(samples, labels)
assert index.backend == "numba"

NumPy remains the default. Numba is supported by CH, WB, DB, XB, GD43, GD53, PS, and cSIL. CONN and rCIP reject backend="numba". CONN’s numerical backend is separate from model_type. Selection is per object and remains fixed through removal of all samples and reinitialization. The batch, incremental, remove, and merge APIs are unchanged.

Compiled kernels accelerate grouping, compactness, centroid distances, and cSIL batch dissimilarities. Dtype-sensitive means and raw-moment reductions remain in NumPy. Unsupported array types (including float16 and non-native byte order) use NumPy for the affected operations. Floating-point rounding may differ; bitwise equivalence is not guaranteed. Undefined criteria return NaN. Unchanged paths, including CH/WB, GD53, and cSIL sample updates and remove/merge operations, are not accelerated. DB, GD43, PS, and XB use compiled centroid-distance kernels during sample updates and after remove/merge. Batch CH/WB compile grouping, compactness, and centroid-to-mean distances; GD53 compiles grouping and compactness. Batch DB, GD43, and XB compile grouping, compactness, and pairwise centroid distances. PS compiles grouping and pairwise centroid distances, without a compactness calculation. cSIL compiles grouping, compactness, and its batch dissimilarity calculation.

The first invocation for a new input type/layout incurs compilation. Later calls reuse compiled code, with a disk cache across processes. Include that initial latency when measuring short tasks, and warm up kernels before measuring steady-state throughput. Numba is imported for the numerical backend only when selected; CONN’s optional ART dependency can independently install and use it. Selecting Numba without the dependency installed raises an installation hint rather than silently selecting another backend.

JAX batch evaluation

Install cvi[jax] or, from a checkout, python -m pip install -e ".[jax]". JAX supports batch CH, WB, and XB:

import jax

jax.config.update("jax_enable_x64", True)
index = cvi.XB(backend="jax")
value = index.get_cvi(samples, labels)

Alternatively, set JAX_ENABLE_X64=1 before starting Python. CVI never changes JAX’s global configuration and reports an error if 64-bit mode is disabled. NumPy remains the default and does not load JAX. The class adapter accepts real integer/float inputs up to 64 bits and arbitrary integer labels. It validates the batch before changing state and retains first-seen label ordering.

The adapter computes input-dtype means with NumPy to preserve existing float32 and large-offset behavior. Compactness, separation, and scores are computed by JAX, then transferred back to the object’s NumPy state and Python scalar result. This adapter synchronizes with the device; it is not intended for use inside jax.jit. Without capacity, incremental updates raise NotImplementedError. Remove and merge are unsupported for all JAX objects. Other indices do not yet support backend="jax".

For device-resident calculations, use the functional interface:

from functools import partial
import jax.numpy as jnp
from cvi.jax import batch_cvi, batch_state, evaluate

x = jnp.array([[0., 0.], [1., 1.], [4., 4.], [5., 5.]])
dense_labels = jnp.array([0, 0, 1, 1])
score = jax.jit(partial(batch_cvi, n_clusters=2, index="XB"))
device_value = score(x, dense_labels)
state = batch_state(x, dense_labels, n_clusters=2)
ch_value = evaluate(state, index="CH")

Functional labels must be dense integers in [0, n_clusters) with every cluster represented. Unlike the class adapter, this interface leaves value-dependent label validation to the caller so it can run under JAX transformations. Invalid partitions produce NaN scores. n_clusters and index must be static under JIT, and new array shapes can recompile. BatchState is an immutable pytree of counts, centroids, centered compactness, global mean, and sample count. It contains JAX arrays and remains on-device.

The functional interface computes in float64, using shifted means to reduce cancellation for large offsets. Reduction order differs from NumPy and may vary by device; bitwise identity is not promised. It supports vmap across batches with compatible shapes and grad with respect to data for fixed labels where the score is differentiable. Undefined CH, WB, and XB scores are NaN when fewer than two clusters are present or their respective denominator is exactly zero: WGSS for CH, BGSS for WB, and minimum centroid separation for XB. These checks use exact zero comparisons, with no epsilon adjustment. Undefined object batch evaluations emit a RuntimeWarning; functional JAX evaluations return NaN without warnings. Time JAX results with block_until_ready() and separate first-use compilation from warmed execution. Host transfers and small CPU workloads can outweigh the benefit of compilation.

Fixed-capacity JAX streaming

Opt into streaming for CH, WB, or XB by reserving cluster slots:

index = cvi.XB(backend="jax", capacity=32)
value = index.get_cvi(samples[0], int(labels[0]))
history = index.update_many(samples[1:100], labels[1:100])
final_value = index.update_many(samples[100:], labels[100:],
                                return_history=False)

capacity is a positive integer limiting the number of distinct labels, not the number of samples. Feature dimension is inferred on the first nonempty call and then fixed. Arbitrary integer labels map to slots in first-seen order. Unused slots do not contribute to the score. Capacity cannot be resized; choose a suitable bound or create a new object. It is only accepted with JAX.

get_cvi accepts a sample, or a one-time initial batch with at least two clusters. An initial batch can be followed by samples or update_many chunks. Chunk processing uses the incremental recurrence in input order; it is distinct from batch initialization. Empty chunks are no-ops and return an empty history or the current score. A score is NaN until at least two clusters are active and the index’s denominator is positive (WGSS for CH, BGSS for WB, or minimum centroid separation for XB). Denominators are checked exactly, without epsilon adjustment. Undefined object batch evaluations warn; incremental and functional streaming calls return NaN without warnings.

The adapter checks the entire input before dispatch: data must be finite real numbers up to 64 bits, labels must be integers, dimensions must match, and all new labels must fit. A rejected sample, batch, or chunk leaves both statistics and label mapping unchanged. A full stream still accepts existing labels. Removal and merging remain unsupported.

Streaming arithmetic is float64, including float32 input conversion. Batch initialization retains the batch adapter’s input-dtype means; subsequent samples use the float64 incremental recurrence, including its residual correction. Floating-point operation ordering can differ across devices, so results are numerically equivalent rather than bitwise identical.

index.stream_state exposes immutable, padded JAX arrays; it is None until the first nonempty call. Numerical state stays on-device. Object methods synchronize to return a Python scalar or NumPy score history. For calculations inside JIT, use the functional API directly:

from cvi.jax import empty_stream, stream_chunk, stream_update

state = empty_stream(capacity=32, n_features=2, index="XB")
state, score = stream_update(state, jnp.array([1., 2.]), 7, index="XB")
state, history = stream_chunk(
    state, jnp.array([[2., 3.], [5., 6.]]), jnp.array([7, 12]), index="XB",
)

Functional labels are slot indices in [0, capacity). Slots may be sparse; this interface does not map external labels. StreamingState contains counts, centroids, compactness, residual corrections, global mean, sample count, active flags, and XB distances. CH/WB use an empty distance array. Use stream_from_batch(batch_state, capacity=32, index="XB") to pad an existing valid BatchState and evaluate_stream(state, index="XB") to evaluate it. CH/WB share a state layout; XB requires its distance-matrix layout.

Functional shape/dtype errors raise ValueError. To work under JIT, out-of-range slot values or nonfinite samples return the unchanged state and NaN output; one invalid row rejects an entire chunk. Check outputs as appropriate for your application. This differs from the object API’s host-side exceptions.

Bind index and return_history statically under JIT. Capacity and feature dimension fix the state shapes, so adding a cluster within capacity does not cause recompilation. Changing chunk length can compile a new specialization. return_history=False avoids allocating a history and evaluates the score only after the scan. CH/WB storage scales as capacity times feature dimension; XB additionally stores a square distance matrix. A large capacity increases work on padded arrays. Prefer chunks for throughput; individual Python calls can cost more than NumPy. See benchmarks/benchmark_jax_stream.py for synchronized compilation, resident chunk, and object timings.