vtx-embed-1M-lf2 / lf2_native.py
Abhaykoul's picture
Initial release of native LF2 2-bit quantized embedding model
3ddb1a3 verified
Raw History Blame Contribute Delete
52.4 kB
"""Vortex-Embed LF2 — Native 2-Bit Embedding Engine.
LF2 format (per weight-block of `block_size` fp32 weights):
- codes: 2-bit levels {0,1,2,3}, 4 weights packed per uint8 byte.
- scale_u8: per-block uint8, double-quantized step
(step = global_scale_max * scale_u8 / 255, step_fp = span / 3).
- min_u8: per-block uint8, double-quantized block minimum
(bmin = global_min + min_u8 / 255 * (global_max - global_min)).
Integer-native guarantee: RAM holds ONLY uint8 bytes
(packed codes + int8 metadata). The only floats in the whole checkpoint
are 3 fp32 scalars per tensor (global_min/max/scale_max, 12 bytes) —
zero FP32/FP16 parameter tables. Dequant happens on-the-fly for the
unique token IDs of each encode batch (registers/L1 temp buffer only),
mirroring the LF4 native engine's pooling/normalize path exactly.
Realtime paths (research: Model2Vec static-lookup + mean-pool O(n*d);
SwiftEmbed SIMD/prefetch/zero-copy; QuIP#/QTIP L1-resident codebooks +
bandwidth-bound fused dequant; AQLM additive LUTs):
- fused numba dequant+mean-pool (no (N,dim) temp, no np.unique, no
torch construction) — default fast path for index loops.
- opt-in preloaded fp32 table (Model2Vec/SwiftEmbed row-index mode).
- opt-in precomputed fp32 step/min meta (skip double-quant per batch).
- truncated-dim early exit (matryoshka needs only leading blocks).
- streaming indexer (tokenize-once, chunked, parallel).
"""
from __future__ import annotations
import json
import os
from pathlib import Path
from typing import Iterator, List, Optional, Sequence, Union
import numpy as np
from safetensors.numpy import load_file, save_file
try:
from tokenizers import Tokenizer
except ImportError: # pragma: no cover
Tokenizer = None
try:
import numba
_NUMBA_OK = True
except ImportError: # pragma: no cover
numba = None # type: ignore
_NUMBA_OK = False
LEVELS = 3 # 2-bit -> levels {0,1,2,3}, step = span / 3
VALS_PER_BYTE = 4
# 256-entry byte->4x2bit LUT (1KB, L1-resident ala QuIP# E8P codebook).
_LUT4 = np.empty((256, VALS_PER_BYTE), dtype=np.uint8)
for _b in range(256):
_LUT4[_b, 0] = (_b >> 0) & 0x03
_LUT4[_b, 1] = (_b >> 2) & 0x03
_LUT4[_b, 2] = (_b >> 4) & 0x03
_LUT4[_b, 3] = (_b >> 6) & 0x03
if _NUMBA_OK:
@numba.njit(cache=True, fastmath=True)
def _fused_seq_full(
packed, scale_u8, min_u8, flat, starts, out,
gmin, grange, smax, nb, bs,
):
# Full-dim specialization: no dd>=dim branch (out_dim == nb*bs).
n = starts.shape[0] - 1
pb = bs // 4
inv255 = 1.0 / 255.0
for d in range(n):
s0 = starts[d]
s1 = starts[d + 1]
L = s1 - s0
if L <= 0:
continue
inv = 1.0 / L
for ti in range(s0, s1):
tid = flat[ti]
dd = 0
for b in range(nb):
step = smax * (scale_u8[tid, b] * inv255)
bmin = gmin + (min_u8[tid, b] * inv255) * grange
poff = b * pb
for k in range(bs):
byte = packed[tid, poff + (k >> 2)]
code = (byte >> ((k & 3) * 2)) & 3
out[d, dd] += (code * step + bmin) * inv
dd += 1
@numba.njit(cache=True, fastmath=True)
def _fused_pool_seq(
packed, scale_u8, min_u8, flat, starts, out,
gmin, grange, smax, dim, nb, bs,
):
n = starts.shape[0] - 1
pb = bs // 4
inv255 = 1.0 / 255.0
for d in range(n):
s0 = starts[d]
s1 = starts[d + 1]
L = s1 - s0
if L <= 0:
continue
inv = 1.0 / L
for ti in range(s0, s1):
tid = flat[ti]
base = 0
for b in range(nb):
step = smax * (scale_u8[tid, b] * inv255)
bmin = gmin + (min_u8[tid, b] * inv255) * grange
poff = b * pb
for k in range(bs):
dd = base + k
if dd >= dim:
break
byte = packed[tid, poff + (k >> 2)]
code = (byte >> ((k & 3) * 2)) & 3
out[d, dd] += (code * step + bmin) * inv
base += bs
@numba.njit(cache=True, fastmath=True, parallel=True)
def _fused_pool_nb(
packed, scale_u8, min_u8, flat, starts, out,
gmin, grange, smax, dim, nb, bs,
):
n = starts.shape[0] - 1
pb = bs // 4
inv255 = 1.0 / 255.0
for d in numba.prange(n):
s0 = starts[d]
s1 = starts[d + 1]
L = s1 - s0
if L <= 0:
continue
inv = 1.0 / L
for ti in range(s0, s1):
tid = flat[ti]
base = 0
for b in range(nb):
step = smax * (scale_u8[tid, b] * inv255)
bmin = gmin + (min_u8[tid, b] * inv255) * grange
poff = b * pb
for k in range(bs):
dd = base + k
if dd >= dim:
break
byte = packed[tid, poff + (k >> 2)]
code = (byte >> ((k & 3) * 2)) & 3
out[d, dd] += (code * step + bmin) * inv
base += bs
@numba.njit(cache=True, fastmath=True)
def _fused_pool_w_seq(
packed, scale_u8, min_u8, flat, starts, out, wrow,
gmin, grange, smax, dim, nb, bs,
):
n = starts.shape[0] - 1
pb = bs // 4
inv255 = 1.0 / 255.0
for d in range(n):
s0 = starts[d]
s1 = starts[d + 1]
wsum = 0.0
for ti in range(s0, s1):
wsum += wrow[ti]
if wsum < 1e-12:
continue
inv = 1.0 / wsum
for ti in range(s0, s1):
tid = flat[ti]
w = wrow[ti] * inv
base = 0
for b in range(nb):
step = smax * (scale_u8[tid, b] * inv255)
bmin = gmin + (min_u8[tid, b] * inv255) * grange
poff = b * pb
for k in range(bs):
dd = base + k
if dd >= dim:
break
byte = packed[tid, poff + (k >> 2)]
code = (byte >> ((k & 3) * 2)) & 3
out[d, dd] += (code * step + bmin) * w
base += bs
@numba.njit(cache=True, fastmath=True, parallel=True)
def _fused_pool_w_nb(
packed, scale_u8, min_u8, flat, starts, out, wrow,
gmin, grange, smax, dim, nb, bs,
):
n = starts.shape[0] - 1
pb = bs // 4
inv255 = 1.0 / 255.0
for d in numba.prange(n):
s0 = starts[d]
s1 = starts[d + 1]
wsum = 0.0
for ti in range(s0, s1):
wsum += wrow[ti]
if wsum < 1e-12:
continue
inv = 1.0 / wsum
for ti in range(s0, s1):
tid = flat[ti]
w = wrow[ti] * inv
base = 0
for b in range(nb):
step = smax * (scale_u8[tid, b] * inv255)
bmin = gmin + (min_u8[tid, b] * inv255) * grange
poff = b * pb
for k in range(bs):
dd = base + k
if dd >= dim:
break
byte = packed[tid, poff + (k >> 2)]
code = (byte >> ((k & 3) * 2)) & 3
out[d, dd] += (code * step + bmin) * w
base += bs
@numba.njit(cache=True, fastmath=True)
def _table_pool_seq(table, flat, starts, out):
n = starts.shape[0] - 1
dim = out.shape[1]
for d in range(n):
s0 = starts[d]
s1 = starts[d + 1]
L = s1 - s0
if L <= 0:
continue
inv = 1.0 / L
for ti in range(s0, s1):
tid = flat[ti]
for j in range(dim):
out[d, j] += table[tid, j] * inv
@numba.njit(cache=True, fastmath=True, parallel=True)
def _table_pool_nb(table, flat, starts, out):
n = starts.shape[0] - 1
dim = out.shape[1]
for d in numba.prange(n):
s0 = starts[d]
s1 = starts[d + 1]
L = s1 - s0
if L <= 0:
continue
inv = 1.0 / L
for ti in range(s0, s1):
tid = flat[ti]
for j in range(dim):
out[d, j] += table[tid, j] * inv
@numba.njit(cache=True, fastmath=True)
def _norm_seq(x):
n = x.shape[0]
dim = x.shape[1]
for i in range(n):
s = 0.0
for j in range(dim):
s += x[i, j] * x[i, j]
inv = 1.0 / (np.sqrt(s) + 1e-12)
for j in range(dim):
x[i, j] *= inv
@numba.njit(cache=True, fastmath=True, parallel=True)
def _norm_nb(x):
n = x.shape[0]
dim = x.shape[1]
for i in numba.prange(n):
s = 0.0
for j in range(dim):
s += x[i, j] * x[i, j]
inv = 1.0 / (np.sqrt(s) + 1e-12)
for j in range(dim):
x[i, j] *= inv
# Parallel crossover: prange thread-spawn costs ~ms cold but the pool
# stays hot after warm_kernels (which warms with a realistic-size
# batch); above this many docs parallel wins by 4-6x. Measured:
# n=256: seq 1.79ms vs par 0.30ms; n=1379: seq 14.3ms vs par 2.9ms.
_PAR_MIN_DOCS = 128
_PAR_MIN_DOCS_NORM = 1024
else: # pragma: no cover
_fused_pool_nb = None
_fused_pool_seq = None
_fused_seq_full = None
_fused_pool_w_nb = None
_fused_pool_w_seq = None
_table_pool_nb = None
_table_pool_seq = None
_norm_nb = None
_norm_seq = None
_PAR_MIN_DOCS = 10**9
_PAR_MIN_DOCS_NORM = 10**9
def quantize_lf2_block(
x: np.ndarray, block_size: int
) -> tuple[np.ndarray, np.ndarray, np.ndarray, float, float, float]:
"""Quantize fp32 matrix to LF2 integer-native format.
Returns (packed_uint8, scale_u8, min_u8, global_min, global_max,
global_scale_max).
"""
x = np.asarray(x, dtype=np.float32)
n, d = x.shape
assert d % block_size == 0, f"dim {d} not divisible by block {block_size}"
n_blocks = d // block_size
xb = x.reshape(n, n_blocks, block_size)
bmin = xb.min(axis=2)
bmax = xb.max(axis=2)
span = (bmax - bmin) / float(LEVELS)
span = np.where(span == 0, 1.0, span)
q = np.clip(np.round((xb - bmin[:, :, None]) / span[:, :, None]), 0, LEVELS)
q = q.astype(np.uint8).reshape(n, d)
step_fp = span # (n, n_blocks) float32
gmin = float(x.min())
gmax = float(x.max())
grange = (gmax - gmin) if gmax > gmin else 1.0
min_u8 = np.clip(
np.round(255.0 * (bmin - gmin) / grange), 0, 255
).astype(np.uint8)
smax = float(step_fp.max())
scale_u8 = np.clip(np.round(255.0 * step_fp / smax), 0, 255).astype(np.uint8)
# Pack 4x 2-bit codes per byte, block-aligned (block_size % 4 == 0 required)
assert block_size % VALS_PER_BYTE == 0
qb = q.reshape(n, -1, VALS_PER_BYTE)
shifts = np.array([0, 2, 4, 6], dtype=np.uint8)
packed = np.zeros((n, q.shape[1] // VALS_PER_BYTE), dtype=np.uint8)
for i in range(VALS_PER_BYTE):
packed |= (qb[:, :, i] << shifts[i]).astype(np.uint8)
return packed, scale_u8, min_u8, gmin, gmax, smax
def dequantize_lf2_meta(
scale_u8: np.ndarray,
min_u8: np.ndarray,
gmin: float,
gmax: float,
smax: float,
) -> tuple[np.ndarray, np.ndarray]:
"""Double-quant metadata -> per-block float step/min (temp buffers only)."""
grange = (gmax - gmin) if gmax > gmin else 1.0
step = smax * scale_u8.astype(np.float32) / 255.0
bmin = gmin + min_u8.astype(np.float32) / 255.0 * grange
return step, bmin
class LF2Config:
def __init__(
self,
vocab_size: int = 29528,
embedding_dim: int = 256,
block_size: int = 32,
num_blocks: int = 8,
global_min: float = 0.0,
global_max: float = 0.0,
global_scale_max: float = 1.0,
matryoshka_dim: Optional[int] = None,
**kwargs,
):
self.vocab_size = vocab_size
self.embedding_dim = embedding_dim
self.block_size = block_size
self.num_blocks = num_blocks
self.global_min = global_min
self.global_max = global_max
self.global_scale_max = global_scale_max
self.matryoshka_dim = matryoshka_dim
@classmethod
def from_dict(cls, d: dict) -> "LF2Config":
return cls(**d)
def to_dict(self) -> dict:
return {
"vocab_size": self.vocab_size,
"embedding_dim": self.embedding_dim,
"block_size": self.block_size,
"num_blocks": self.num_blocks,
"global_min": self.global_min,
"global_max": self.global_max,
"global_scale_max": self.global_scale_max,
"matryoshka_dim": self.matryoshka_dim,
"quantization": "lf2",
"bits": 2,
}
class VortexEmbedLF2:
"""Native 2-bit sentence embedding model. Same encode path as LF4 engine."""
def __init__(
self,
packed: np.ndarray,
scale_u8: np.ndarray,
min_u8: np.ndarray,
tokenizer_data: Union[str, Path],
config: Union[dict, LF2Config],
*,
matryoshka_dim: Optional[int] = None,
) -> None:
self.packed = np.asarray(packed, dtype=np.uint8)
self.scale_u8 = np.asarray(scale_u8, dtype=np.uint8)
self.min_u8 = np.asarray(min_u8, dtype=np.uint8)
self.tokenizer_data = str(tokenizer_data)
self.config = (
config if isinstance(config, LF2Config) else LF2Config.from_dict(config)
)
self.vocab_size = int(self.config.vocab_size)
self.dim = int(self.config.embedding_dim)
self.block_size = int(self.config.block_size)
self.num_blocks = int(self.config.num_blocks)
self.matryoshka_dim = matryoshka_dim or self.config.matryoshka_dim
self._tokenizer: Optional[Tokenizer] = None
self._sif_weights: Optional[np.ndarray] = None
self._pc_directions: Optional[np.ndarray] = None
# SIF 'a' mirrors the LF4 single-file engine default.
self.sif_a: float = 0.05
self.sif_pc: float = 1.0
self.pc_k: int = 1
# Hot-token row cache (opt-in, runtime only — not a parameter table).
# Maps token id -> dequantized fp32 row. Disabled by default to keep
# the native guarantee; enable with enable_cache() for indexing loops
# over skewed corpora. Not shared across threads (use shallow_clone).
self._row_cache: Optional[dict] = None
self._row_cache_cap: int = 0
# Opt-in runtime accelerators (not parameter tables; transient temp).
# preload_meta(): fp32 step/min per (vocab, block) — skips per-batch
# double-quant (Model2Vec-style precompute, ~2x240KB for 30k vocab).
# preload_table(): full fp32 table (V,D) — SwiftEmbed row-index mode,
# max tok/s at ~30MB transient (freed with unload_table()).
self._step_f32: Optional[np.ndarray] = None
self._bmin_f32: Optional[np.ndarray] = None
self._fp_table: Optional[np.ndarray] = None
# Warm numba kernels at construction (compile once, not per batch).
self._nb_warmed: bool = False
def enable_cache(self, cap: int = 4096) -> "VortexEmbedLF2":
"""Opt-in LRU cache of dequantized token rows (runtime temp only)."""
from collections import OrderedDict
self._row_cache = OrderedDict()
self._row_cache_cap = max(int(cap), 1)
return self
def disable_cache(self) -> "VortexEmbedLF2":
self._row_cache = None
self._row_cache_cap = 0
return self
def shallow_clone(self) -> "VortexEmbedLF2":
"""Share read-only params, fresh fit-state and empty cache config."""
c = VortexEmbedLF2(
self.packed, self.scale_u8, self.min_u8,
self.tokenizer_data, self.config,
matryoshka_dim=self.matryoshka_dim,
)
c.sif_a, c.sif_pc, c.pc_k = self.sif_a, self.sif_pc, self.pc_k
if self._row_cache is not None:
c.enable_cache(self._row_cache_cap)
# Share accelerators read-only (they are deterministic of params).
c._step_f32, c._bmin_f32, c._fp_table = (
self._step_f32, self._bmin_f32, self._fp_table)
return c
# -- realtime accelerators (opt-in, runtime temp only) ---------------
def preload_meta(self) -> "VortexEmbedLF2":
"""Precompute fp32 step/min tables (skips per-batch double-quant)."""
step, bmin = dequantize_lf2_meta(
np.arange(256, dtype=np.uint8)[self.scale_u8.ravel()].reshape(
self.scale_u8.shape) * 0 + self.scale_u8,
self.min_u8,
self.config.global_min, self.config.global_max,
self.config.global_scale_max,
) if False else dequantize_lf2_meta(
self.scale_u8, self.min_u8,
self.config.global_min, self.config.global_max,
self.config.global_scale_max,
)
self._step_f32 = np.ascontiguousarray(step, dtype=np.float32)
self._bmin_f32 = np.ascontiguousarray(bmin, dtype=np.float32)
return self
def unload_meta(self) -> "VortexEmbedLF2":
self._step_f32 = None
self._bmin_f32 = None
return self
def preload_table(self, dim: Optional[int] = None) -> "VortexEmbedLF2":
"""Materialize full fp32 table transiently (SwiftEmbed row-index mode).
`dim` truncates columns (matryoshka early-exit). Call unload_table()
to restore the integer-native guarantee.
"""
d = dim or self.dim
full = self._dequantize_fresh(
np.arange(self.vocab_size, dtype=np.int64))[:, :d]
self._fp_table = np.ascontiguousarray(full, dtype=np.float32)
return self
def unload_table(self) -> "VortexEmbedLF2":
self._fp_table = None
return self
def warm_kernels(self) -> "VortexEmbedLF2":
"""Compile numba kernels once (avoid first-batch compile stall)."""
if not _NUMBA_OK or self._nb_warmed:
return self
try:
pk = np.ascontiguousarray(self.packed[:8])
sc = np.ascontiguousarray(self.scale_u8[:8])
mn = np.ascontiguousarray(self.min_u8[:8])
gmin = float(self.config.global_min)
gr = float((self.config.global_max - self.config.global_min)
if self.config.global_max > self.config.global_min else 1.0)
smax = float(self.config.global_scale_max)
nb, bs = self.num_blocks, self.block_size
# Warm BOTH seq + par variants (dispatch picks by batch size).
# The par warm uses a realistic-size batch so the numba thread
# pool is hot before serving (cold spawn costs ~ms).
nW = 512
flW = np.random.default_rng(0).integers(
0, 8, size=2048).astype(np.int64)
stW = np.linspace(0, 2048, nW + 1).astype(np.int64)
ouW = np.zeros((nW, self.dim), dtype=np.float32)
_fused_pool_nb(pk, sc, mn, flW, stW, ouW, gmin, gr, smax,
self.dim, nb, bs)
fl = np.array([0, 1, 2], dtype=np.int64)
st = np.array([0, 2, 3], dtype=np.int64)
ou = np.zeros((2, self.dim), dtype=np.float32)
_fused_seq_full(pk, sc, mn, fl, st, ou, gmin, gr, smax, nb, bs)
ou[:] = 0
_fused_pool_seq(pk, sc, mn, fl, st, ou, gmin, gr, smax,
self.dim, nb, bs)
_norm_seq(ou)
_norm_nb(ou)
self._nb_warmed = True
except Exception:
pass
return self
# -- properties -----------------------------------------------------
@property
def tokenizer(self) -> Tokenizer:
if self._tokenizer is None:
if Tokenizer is None: # pragma: no cover
raise RuntimeError("tokenizers required: pip install tokenizers")
self._tokenizer = Tokenizer.from_file(self.tokenizer_data)
return self._tokenizer
@property
def int_bytes(self) -> int:
"""Integer parameter bytes in RAM (codes + int8 meta)."""
return (
int(self.packed.nbytes)
+ int(self.scale_u8.nbytes)
+ int(self.min_u8.nbytes)
)
@property
def model_size_mb(self) -> float:
return (self.int_bytes + 12) / 1e6 # +3 fp32 global scalars
@property
def on_disk_size_mb(self) -> float:
return (self.int_bytes + 12) / 1e6
# -- io --------------------------------------------------------------
@classmethod
def quantize_from_matrix(
cls,
w_fp32: np.ndarray,
tokenizer_data: Union[str, Path],
block_size: int = 32,
matryoshka_dim: Optional[int] = None,
) -> "VortexEmbedLF2":
packed, scale_u8, min_u8, gmin, gmax, smax = quantize_lf2_block(
w_fp32, block_size
)
n, d = w_fp32.shape
cfg = LF2Config(
vocab_size=n,
embedding_dim=d,
block_size=block_size,
num_blocks=d // block_size,
global_min=gmin,
global_max=gmax,
global_scale_max=smax,
matryoshka_dim=matryoshka_dim,
)
return cls(packed, scale_u8, min_u8, tokenizer_data, cfg,
matryoshka_dim=matryoshka_dim)
def save_pretrained(self, path: Union[str, Path]) -> None:
out = Path(path)
out.mkdir(parents=True, exist_ok=True)
save_file(
{
"embedding_packed": self.packed,
"embedding_scale_u8": self.scale_u8,
"embedding_min_u8": self.min_u8,
},
str(out / "model.safetensors"),
)
(out / "config.json").write_text(json.dumps(self.config.to_dict(), indent=2))
if not (out / "tokenizer.json").exists():
(out / "tokenizer.json").write_text(Path(self.tokenizer_data).read_text())
@classmethod
def from_pretrained(
cls, path: Union[str, Path], matryoshka_dim: Optional[int] = None
) -> "VortexEmbedLF2":
path = Path(path)
tensors = load_file(str(path / "model.safetensors"))
config = json.loads((path / "config.json").read_text())
return cls(
packed=tensors["embedding_packed"],
scale_u8=tensors["embedding_scale_u8"],
min_u8=tensors["embedding_min_u8"],
tokenizer_data=str(path / "tokenizer.json"),
config=config,
matryoshka_dim=matryoshka_dim,
)
def dequantize_all(self) -> np.ndarray:
"""Full fp32 table (offline analysis only — never held in RAM at runtime)."""
return self.dequantize_ids(
np.arange(self.vocab_size, dtype=np.int64))
# -- SIF-IDF + PC removal (mirrors LF4 engine) -------------------------
def fit_idf(self, corpus_token_lists: Sequence[Sequence[int]]) -> "VortexEmbedLF2":
flat = (
np.concatenate(corpus_token_lists)
if corpus_token_lists
else np.empty(0, dtype=np.int64)
)
total = max(int(flat.size), 1)
counts = np.bincount(flat, minlength=self.vocab_size).astype(np.float64)
p = counts / total
denom = self.sif_a + p
with np.errstate(divide="ignore", invalid="ignore"):
weights = np.where(p > 0, self.sif_a / denom, 1.0)
self._sif_weights = weights.astype(np.float32)
return self
def fit_pc(
self, corpus_embeddings: np.ndarray, k: Optional[int] = None
) -> "VortexEmbedLF2":
if k is None:
k = self.pc_k
if corpus_embeddings.size == 0 or k <= 0:
return self
x = corpus_embeddings.astype(np.float32)
x = x - x.mean(axis=0, keepdims=True)
try:
_, _, vt = np.linalg.svd(x, full_matrices=False)
pcs = vt[:k].astype(np.float32)
pcs = pcs / (np.linalg.norm(pcs, axis=1, keepdims=True) + 1e-12)
self._pc_directions = pcs
except np.linalg.LinAlgError:
self._pc_directions = None
return self
def _apply_pc(self, x: np.ndarray) -> np.ndarray:
if self.sif_pc <= 0 or self._pc_directions is None:
return x
out = x
for pc in self._pc_directions:
proj = (out @ pc)[:, None] * pc[None, :]
out = out - self.sif_pc * proj
return out
def reset_fit(self) -> "VortexEmbedLF2":
self._sif_weights = None
self._pc_directions = None
return self
# -- native on-the-fly dequant (H17: planar-fill + single astype) ----
def dequantize_ids(self, token_ids: np.ndarray) -> np.ndarray:
"""2-pass dequant: strided fills happen on a uint8 temp (1/4 the
traffic), then ONE u8->f32 cast + ONE contiguous blocked fmadd.
No (N, dim) float temp, no per-stream astype."""
if token_ids.size == 0:
return np.empty((0, self.dim), dtype=np.float32)
n = len(token_ids)
nb, bs = self.num_blocks, self.block_size
cache = self._row_cache
if cache is not None and n <= 512:
# Hot-token path: reuse cached rows, dequantize misses only.
rows: List[Optional[np.ndarray]] = [cache.get(int(t)) for t in token_ids]
# LRU touch on hits
for t, r in zip(token_ids, rows):
if r is not None:
cache.move_to_end(int(t))
missing = np.array(
[t for t, r in zip(token_ids, rows) if r is None], dtype=np.int64
)
if missing.size:
got = self._dequantize_fresh(missing)
for t, r in zip(missing, got):
cache[int(t)] = r
if len(cache) > self._row_cache_cap:
cache.popitem(last=False)
it = iter(zip(missing, got))
lut = {int(t): r for t, r in it}
rows = [r if r is not None else lut[int(t)]
for t, r in zip(token_ids, rows)]
return np.stack(list(rows)).astype(np.float32)
return self._dequantize_fresh(np.asarray(token_ids, dtype=np.int64))
def _dequantize_fresh(self, token_ids: np.ndarray,
out_dim: Optional[int] = None) -> np.ndarray:
"""2-pass dequant: strided fills happen on a uint8 temp (1/4 the
traffic), then ONE u8->f32 cast + ONE contiguous blocked fmadd.
No (N, dim) float temp, no per-stream astype.
`out_dim` enables matryoshka early-exit (leading blocks only).
Uses preloaded fp32 meta when available (skips double-quant).
"""
if token_ids.size == 0:
d = out_dim or self.dim
return np.empty((0, d), dtype=np.float32)
n = len(token_ids)
nb, bs = self.num_blocks, self.block_size
dim = out_dim or self.dim
nb_need = min(nb, (dim + bs - 1) // bs)
p = self.packed[token_ids]
if nb_need < nb:
# Slice leading bytes/blocks only (truncated-dim early exit).
pb = bs // VALS_PER_BYTE
p = p[:, : nb_need * pb]
p = p.reshape(n, nb_need, bs // VALS_PER_BYTE)
t = np.empty((n, nb_need, bs), dtype=np.uint8)
t[:, :, 0::4] = p & 0x03
t[:, :, 1::4] = (p >> 2) & 0x03
t[:, :, 2::4] = (p >> 4) & 0x03
t[:, :, 3::4] = (p >> 6) & 0x03
if self._step_f32 is not None and self._bmin_f32 is not None:
step = self._step_f32[token_ids][:, :nb_need]
bmin = self._bmin_f32[token_ids][:, :nb_need]
else:
step, bmin = dequantize_lf2_meta(
self.scale_u8[token_ids][:, :nb_need]
if nb_need < nb else self.scale_u8[token_ids],
self.min_u8[token_ids][:, :nb_need]
if nb_need < nb else self.min_u8[token_ids],
self.config.global_min,
self.config.global_max,
self.config.global_scale_max,
)
f = t.astype(np.float32)
f *= step[:, :, None]
f += bmin[:, :, None]
return f.reshape(n, nb_need * bs)[:, :dim]
def _dequantize_lut(self, token_ids: np.ndarray,
out_dim: Optional[int] = None) -> np.ndarray:
"""LUT-gather variant (256x4 table, QuIP#-style L1 codebook)."""
if token_ids.size == 0:
return np.empty((0, out_dim or self.dim), dtype=np.float32)
n = len(token_ids)
nb, bs = self.num_blocks, self.block_size
dim = out_dim or self.dim
nb_need = min(nb, (dim + bs - 1) // bs)
pb = bs // VALS_PER_BYTE
p = self.packed[token_ids]
if nb_need < nb:
p = p[:, : nb_need * pb]
t = _LUT4[p] # (n, bytes, 4) uint8, single gather
t = t.reshape(n, nb_need, bs)
if self._step_f32 is not None and self._bmin_f32 is not None:
step = self._step_f32[token_ids][:, :nb_need]
bmin = self._bmin_f32[token_ids][:, :nb_need]
else:
step, bmin = dequantize_lf2_meta(
self.scale_u8[token_ids][:, :nb_need]
if nb_need < nb else self.scale_u8[token_ids],
self.min_u8[token_ids][:, :nb_need]
if nb_need < nb else self.min_u8[token_ids],
self.config.global_min,
self.config.global_max,
self.config.global_scale_max,
)
f = t.astype(np.float32)
f *= step[:, :, None]
f += bmin[:, :, None]
return f.reshape(n, nb_need * bs)[:, :dim]
# -- encode (same segment-sum path as LF4 engine) ----------------------
def _tokenize_batch(self, texts: Sequence[str]) -> List[List[int]]:
encoded = self.tokenizer.encode_batch(list(texts))
return [
[tid for tid in item.ids if 0 <= int(tid) < self.vocab_size]
for item in encoded
]
@staticmethod
def _normalize_inplace(x: np.ndarray) -> None:
norms = np.linalg.norm(x, axis=1, keepdims=True)
np.divide(x, np.maximum(norms, 1e-12), out=x)
@staticmethod
def _flat_starts(token_lists: Sequence[Sequence[int]],
max_tokens: int = 0):
n = len(token_lists)
if n == 0:
return (np.empty(0, dtype=np.int64), np.zeros(1, dtype=np.int64),
np.empty(0, dtype=np.int64))
if max_tokens and max_tokens > 0:
trunc = [ids[:max_tokens] if len(ids) > max_tokens else ids
for ids in token_lists]
else:
trunc = list(token_lists)
# One Python-level pass with C-speed list.extend + a single
# array build (2x faster than np.concatenate's per-list convert).
big: List[int] = []
ap = big.extend
for t in trunc:
ap(t)
lens = np.fromiter((len(t) for t in trunc), dtype=np.int64, count=n)
if big:
flat = np.asarray(big, dtype=np.int64)
else:
flat = np.empty(0, dtype=np.int64)
starts = np.empty(n + 1, dtype=np.int64)
starts[0] = 0
np.cumsum(lens, out=starts[1:])
return flat, starts, lens
def _encode_fused(self, token_lists, *, normalize: bool,
out_dim: int, max_tokens: int = 0) -> Optional[np.ndarray]:
"""Fused numba dequant+pool: no unique, no (T,dim) temp, no torch."""
if not _NUMBA_OK or self._pc_directions is not None:
return None
flat, starts, lens = self._flat_starts(token_lists, max_tokens)
n = len(token_lists)
if flat.size == 0:
return np.zeros((n, out_dim), dtype=np.float32)
out = np.zeros((n, out_dim), dtype=np.float32)
nb_need = min(self.num_blocks, (out_dim + self.block_size - 1) // self.block_size)
gmin = float(self.config.global_min)
gmax = float(self.config.global_max)
grange = (gmax - gmin) if gmax > gmin else 1.0
smax = float(self.config.global_scale_max)
par = n >= _PAR_MIN_DOCS
try:
if self._sif_weights is not None:
wrow = self._sif_weights[flat].astype(np.float32)
kern = _fused_pool_w_nb if par else _fused_pool_w_seq
kern(self.packed, self.scale_u8, self.min_u8,
flat, starts, out, wrow,
gmin, grange, smax, out_dim,
nb_need, self.block_size)
elif out_dim == nb_need * self.block_size and not par:
# Fast lane: full-dim seq kernel, no bounds branch.
_fused_seq_full(self.packed, self.scale_u8, self.min_u8,
flat, starts, out,
gmin, grange, smax,
nb_need, self.block_size)
else:
kern = _fused_pool_nb if par else _fused_pool_seq
kern(self.packed, self.scale_u8, self.min_u8,
flat, starts, out,
gmin, grange, smax, out_dim,
nb_need, self.block_size)
except Exception:
return None
if normalize and n:
if _norm_nb is not None:
try:
(_norm_nb if n >= _PAR_MIN_DOCS_NORM else _norm_seq)(out)
except Exception:
self._normalize_inplace(out)
else:
self._normalize_inplace(out)
return out
def _encode_table(self, token_lists, *, normalize: bool,
out_dim: int, max_tokens: int = 0) -> Optional[np.ndarray]:
"""Preloaded-fp32 row-index pool (Model2Vec/SwiftEmbed mode)."""
if self._fp_table is None:
return None
if self._pc_directions is not None:
return None
tab = self._fp_table
if tab.shape[1] < out_dim:
return None
tab = tab[:, :out_dim]
flat, starts, _ = self._flat_starts(token_lists, max_tokens)
n = len(token_lists)
if flat.size == 0:
return np.zeros((n, out_dim), dtype=np.float32)
out = np.zeros((n, out_dim), dtype=np.float32)
try:
if self._sif_weights is not None or not _NUMBA_OK:
raise RuntimeError("fallback")
kern = _table_pool_nb if n >= _PAR_MIN_DOCS else _table_pool_seq
kern(np.ascontiguousarray(tab), flat, starts, out)
except Exception:
# Numpy fallback: unique-dedup gather + reduceat (still no dequant).
uq, inv = np.unique(flat, return_inverse=True)
te = np.ascontiguousarray(tab[uq])[inv]
if self._sif_weights is not None:
w = self._sif_weights[flat].astype(np.float32)[:, None]
te = te * w
ends = starts[1:]
bounds = starts[:-1]
# guard empty docs: reduceat needs valid indices; handle via mask
sums = np.add.reduceat(te, bounds, axis=0) if te.size else out
lens = np.diff(starts).astype(np.float32)
if self._sif_weights is not None:
wf = self._sif_weights[flat].astype(np.float32)
wpr = np.add.reduceat(wf, bounds)
wpr = np.maximum(wpr, 1e-12)
else:
wpr = np.maximum(lens, 1.0)
# Fix rows for empty docs (reduceat wraps around): zero them.
out = sums / wpr[:, None]
out[lens == 0] = 0.0
if normalize and n:
self._normalize_inplace(out)
return out.astype(np.float32)
if normalize and n:
self._normalize_inplace(out)
return out
def _encode_subbatch(
self, token_lists: Sequence[Sequence[int]], *, normalize: bool,
fast: bool = True,
) -> np.ndarray:
n = len(token_lists)
if fast:
# Fast dispatch order: table (fastest) -> fused numba (no temp)
# -> legacy unique+torch (exact legacy numerics, SIF/PC-safe).
got = self._encode_table(token_lists, normalize=False,
out_dim=self.dim)
if got is not None:
embs = self._apply_pc(got)
if normalize:
self._normalize_inplace(embs)
return embs
got = self._encode_fused(token_lists, normalize=False,
out_dim=self.dim)
if got is not None:
embs = self._apply_pc(got)
if normalize:
self._normalize_inplace(embs)
return embs
return self._encode_legacy(token_lists, normalize=normalize)
def _encode_legacy(self, token_lists, *, normalize: bool) -> np.ndarray:
n = len(token_lists)
flat = (
np.concatenate(token_lists)
if token_lists
else np.empty(0, dtype=np.int64)
)
if flat.size == 0:
return np.zeros((n, self.dim), dtype=np.float32)
unique_ids, inverse = np.unique(flat, return_inverse=True)
token_embs = self.dequantize_ids(unique_ids)[inverse]
if self._sif_weights is not None:
w = self._sif_weights[flat].astype(np.float32)[:, None]
token_embs = token_embs * w
try:
import torch
ro = torch.from_numpy(
np.repeat(
np.arange(n, dtype=np.int64),
[len(ids) for ids in token_lists],
)
)
em = torch.from_numpy(np.ascontiguousarray(token_embs))
sums = torch.zeros((n, self.dim), dtype=torch.float32)
sums.index_add_(0, ro, em)
sums = sums.numpy()
except ImportError:
chunk_lens = np.array(
[len(ids) for ids in token_lists], dtype=np.int64
)
ends = np.cumsum(chunk_lens)
bounds = np.empty(n + 1, dtype=np.int64)
bounds[0] = 0
bounds[1:] = ends
sums = np.add.reduceat(token_embs, bounds[:-1], axis=0)
chunk_lens = np.array([len(ids) for ids in token_lists], dtype=np.int64)
if self._sif_weights is not None:
w_full = self._sif_weights[flat].astype(np.float32)
ends = np.cumsum(chunk_lens)
bounds = np.empty(n + 1, dtype=np.int64)
bounds[0] = 0
bounds[1:] = ends
w_per_row = np.add.reduceat(w_full, bounds[:-1])
w_per_row = np.maximum(w_per_row, 1e-12)
else:
w_per_row = np.maximum(chunk_lens.astype(np.float32), 1.0)
embs = sums / w_per_row[:, None]
embs = self._apply_pc(embs)
if normalize:
self._normalize_inplace(embs)
return embs
def encode_batch(
self,
texts: Sequence[str],
*,
normalize: bool = True,
truncate_dim: Optional[int] = None,
) -> np.ndarray:
if not texts:
return np.zeros((0, self.dim), dtype=np.float32)
if len(texts) == 1:
# H17 latency path: single text needs no segment sum — plain
# (weighted) mean skips torch construction + index_add entirely.
embs = self._encode_single(texts[0])
embs = embs[None, :]
else:
embs = self._encode_subbatch(self._tokenize_batch(list(texts)),
normalize=False)
dim = truncate_dim if truncate_dim is not None else self.matryoshka_dim
if dim is not None and 0 < dim < self.dim:
embs = embs[:, :dim]
if normalize and embs.shape[0] > 0:
self._normalize_inplace(embs)
return embs
def _encode_single(self, text: str) -> np.ndarray:
ids = self._tokenize_batch([text])[0]
if not ids:
return np.zeros((self.dim,), dtype=np.float32)
flat = np.asarray(ids, dtype=np.int64)
if flat.size <= 64:
# Short-text path: mean over occurrences needs no dedup math —
# dequantize flat directly, skipping unique + inverse gather.
# (Identical to dedup-then-mean up to fp summation order.)
token_embs = self.dequantize_ids(flat)
else:
unique_ids, inverse = np.unique(flat, return_inverse=True)
token_embs = self.dequantize_ids(unique_ids)[inverse]
if self._sif_weights is not None:
w = self._sif_weights[flat].astype(np.float32)
embs = (token_embs * w[:, None]).sum(axis=0) / max(float(w.sum()), 1e-12)
else:
embs = token_embs.mean(axis=0)
return self._apply_pc(embs[None, :])[0]
def encode(
self,
texts: Union[str, Sequence[str]],
*,
normalize: bool = True,
truncate_dim: Optional[int] = None,
) -> np.ndarray:
if isinstance(texts, str):
return self.encode_batch(
[texts], normalize=normalize, truncate_dim=truncate_dim
)[0]
return self.encode_batch(
list(texts), normalize=normalize, truncate_dim=truncate_dim
)
# -- realtime indexing API (pre-tokenized + parallel) ------------------
def encode_ids(
self,
token_lists: Sequence[Sequence[int]],
*,
normalize: bool = True,
truncate_dim: Optional[int] = None,
max_tokens: Optional[int] = None,
fast: bool = True,
) -> np.ndarray:
"""Encode pre-tokenized id lists (skips the tokenizer entirely).
`max_tokens` truncates each list (Model2Vec-style max_length).
Reactive indexers tokenize once, then call this per batch.
`fast=True` (default) uses table/fused kernels with truncated-dim
early-exit; `fast=False` forces the legacy unique+torch path.
"""
lists: List[Sequence[int]] = list(token_lists)
mt = int(max_tokens) if max_tokens is not None and max_tokens > 0 else 0
if not lists:
d0 = truncate_dim or self.matryoshka_dim or self.dim
return np.zeros((0, d0), dtype=np.float32)
dim = truncate_dim if truncate_dim is not None else self.matryoshka_dim
out_dim = dim if dim is not None and 0 < dim < self.dim else self.dim
if fast:
got = self._encode_table(lists, normalize=False,
out_dim=out_dim, max_tokens=mt)
if got is None:
got = self._encode_fused(lists, normalize=False,
out_dim=out_dim, max_tokens=mt)
if got is not None:
got = self._apply_pc(got)
if normalize and got.shape[0]:
self._normalize_inplace(got)
return got
if mt:
lists = [ids[:mt] for ids in lists]
if len(lists) == 1:
flat = np.asarray(lists[0], dtype=np.int64)
if flat.size == 0:
embs = np.zeros((1, self.dim), dtype=np.float32)
else:
embs = self._mean_pool(flat)[None, :]
else:
embs = self._encode_subbatch(lists, normalize=False, fast=fast)
if out_dim < self.dim:
embs = embs[:, :out_dim]
if normalize and embs.shape[0] > 0:
self._normalize_inplace(embs)
return embs
def _mean_pool(self, flat: np.ndarray) -> np.ndarray:
if flat.size <= 64:
token_embs = self.dequantize_ids(flat)
else:
unique_ids, inverse = np.unique(flat, return_inverse=True)
token_embs = self.dequantize_ids(unique_ids)[inverse]
if self._sif_weights is not None:
w = self._sif_weights[flat].astype(np.float32)
embs = (token_embs * w[:, None]).sum(axis=0) / max(float(w.sum()), 1e-12)
else:
embs = token_embs.mean(axis=0)
return self._apply_pc(embs[None, :])[0]
def _fast_usable(self) -> bool:
"""True when a single-shot numba path handles this config.
The fused/table kernels already parallelize over docs internally
(prange), so ThreadPool sharding on top only adds spawn + order
restore overhead. Single-shot is the fastest option.
"""
return bool(_NUMBA_OK) and self._pc_directions is None
def encode_parallel(
self,
token_lists: Sequence[Sequence[int]],
*,
n_jobs: int = 8,
batch: int = 256,
normalize: bool = True,
truncate_dim: Optional[int] = None,
max_tokens: Optional[int] = None,
fast: bool = True,
) -> np.ndarray:
"""Shard pre-tokenized id lists across worker threads.
Tokenize ONCE in the caller, then fan out pure vector math.
When the single-shot numba fast path applies (default: no PC fit),
`n_jobs`/`batch` are bypassed — one call already saturates cores.
Threaded sharding remains for the legacy torch path (`fast=False`)
or PC-fitted models.
"""
if fast and self._fast_usable():
return self.encode_ids(
token_lists, normalize=normalize, truncate_dim=truncate_dim,
max_tokens=max_tokens, fast=True)
from concurrent.futures import ThreadPoolExecutor
lists: List[Sequence[int]] = list(token_lists)
if max_tokens is not None and max_tokens > 0:
lists = [ids[:max_tokens] for ids in lists]
n = len(lists)
if n == 0:
d = truncate_dim or self.matryoshka_dim or self.dim
return np.zeros((0, d), dtype=np.float32)
if n_jobs <= 0:
n_jobs = max(1, os.cpu_count() or 1)
n_jobs = max(1, min(int(n_jobs), n))
if n_jobs == 1:
return self.encode_ids(
lists, normalize=normalize, truncate_dim=truncate_dim,
fast=fast,
)
chunks = [lists[i::n_jobs] for i in range(n_jobs)]
workers = [self.shallow_clone() for _ in range(n_jobs)]
def _run(args) -> np.ndarray:
w, ch = args
out = []
for i in range(0, len(ch), batch):
out.append(
w.encode_ids(
ch[i:i + batch],
normalize=False,
truncate_dim=truncate_dim,
fast=fast,
)
)
return np.vstack(out) if out else np.zeros((0, self.dim))
with ThreadPoolExecutor(max_workers=n_jobs) as ex:
parts = list(ex.map(_run, zip(workers, chunks)))
# Restore original order (round-robin interleave invert)
order = np.argsort(
np.concatenate([np.arange(i, n, n_jobs) for i in range(n_jobs)])
)
embs = np.vstack(parts)[order]
dim = truncate_dim if truncate_dim is not None else self.matryoshka_dim
if dim is not None and 0 < dim < self.dim and embs.shape[1] > dim:
embs = embs[:, :dim]
if normalize and embs.shape[0] > 0:
self._normalize_inplace(embs)
return embs
# -- streaming realtime indexer -------------------------------------
def tokenize_texts(self, texts: Sequence[str],
max_tokens: int = 0) -> List[List[int]]:
"""Tokenize once (main thread); reuse lists for encode_ids*."""
lists = self._tokenize_batch(list(texts))
if max_tokens and max_tokens > 0:
lists = [ids[:max_tokens] for ids in lists]
return lists
def index_texts(self, texts: Sequence[str], *,
batch: int = 512, n_jobs: int = 0,
normalize: bool = True,
truncate_dim: Optional[int] = None,
max_tokens: Optional[int] = None,
fast: bool = True,
show_progress: bool = False) -> np.ndarray:
"""End-to-end realtime index: tokenize-once + chunked parallel pool.
Tokenizes the whole input in ONE tokenizer call (Rust-batched),
then pools each 50k-doc chunk in a single numba shot. `n_jobs`/
`batch` only affect the legacy path; the fast path ignores them
(internal prange already saturates cores).
"""
texts = list(texts)
n = len(texts)
d = truncate_dim or self.matryoshka_dim or self.dim
if n == 0:
return np.zeros((0, d), dtype=np.float32)
mt = int(max_tokens or 0)
if n <= 50000:
lists = self.tokenize_texts(texts, max_tokens=mt)
return self.encode_ids(
lists, normalize=normalize, truncate_dim=truncate_dim,
fast=fast)
out_parts: List[np.ndarray] = []
step = 50000
it = range(0, n, step)
if show_progress:
try:
from tqdm import tqdm # type: ignore
it = tqdm(it, desc="index")
except Exception:
pass
for s in it:
lists = self.tokenize_texts(texts[s:s + step], max_tokens=mt)
out_parts.append(self.encode_ids(
lists, normalize=normalize, truncate_dim=truncate_dim,
fast=fast))
return np.vstack(out_parts) if out_parts else np.zeros((0, d))
def index_stream(self, texts: Iterator[str], *,
batch: int = 512, n_jobs: int = 0,
normalize: bool = True,
truncate_dim: Optional[int] = None,
max_tokens: Optional[int] = None,
fast: bool = True) -> Iterator[np.ndarray]:
"""Yield embedding chunks for an unbounded text iterator."""
buf: List[str] = []
width = max(int(batch) * max(int(n_jobs or 1), 1), 512)
for t in texts:
buf.append(t)
if len(buf) >= width:
lists = self.tokenize_texts(buf, int(max_tokens or 0))
yield self.encode_parallel(
lists, n_jobs=n_jobs or 1, batch=int(batch),
normalize=normalize, truncate_dim=truncate_dim,
fast=fast)
buf = []
if buf:
lists = self.tokenize_texts(buf, int(max_tokens or 0))
yield self.encode_parallel(
lists, n_jobs=n_jobs or 1, batch=int(batch),
normalize=normalize, truncate_dim=truncate_dim,
fast=fast)
def search(self, queries: np.ndarray, index: np.ndarray, top_k: int = 10,
index_normalized: bool = False) -> tuple[np.ndarray, np.ndarray]:
"""Cosine top-k (queries assumed L2-normalized)."""
q = np.asarray(queries, dtype=np.float32)
if q.ndim == 1:
q = q[None, :]
idx = index if index_normalized else (
index / np.maximum(np.linalg.norm(index, axis=1, keepdims=True), 1e-12))
sims = q @ idx.T
k = max(1, min(int(top_k), index.shape[0]))
part = np.argpartition(-sims, k - 1, axis=1)[:, :k]
row = np.take_along_axis(sims, part, axis=1)
order = np.argsort(-row, axis=1)
idx_out = np.take_along_axis(part, order, axis=1)
sco_out = np.take_along_axis(row, order, axis=1)
return sco_out, idx_out