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