Download lf2_native.py from VTXAI/vtx-embed-1M-lf2: direct link, hf CLI and curl.
- Browser
- Download file 52.4 kB
-
https://huggingface.co/VTXAI/vtx-embed-1M-lf2/resolve/main/lf2_native.py
- Command line
-
hf download hf://VTXAI/vtx-embed-1M-lf2/lf2_native.py
-
curl -L -o lf2_native.py https://huggingface.co/VTXAI/vtx-embed-1M-lf2/resolve/main/lf2_native.py
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: | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 | |
| 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 ----------------------------------------------------- | |
| 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 | |
| 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) | |
| ) | |
| def model_size_mb(self) -> float: | |
| return (self.int_bytes + 12) / 1e6 # +3 fp32 global scalars | |
| def on_disk_size_mb(self) -> float: | |
| return (self.int_bytes + 12) / 1e6 | |
| # -- io -------------------------------------------------------------- | |
| 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()) | |
| 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 | |
| ] | |
| 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) | |
| 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 | |