codekingpro/portable-devtools
115k
1import os2import sys3import re4import tempfile5import unicodedata6from pathlib import Path7from itertools import islice8from typing import Iterable, TypeVar9 10import numpy as np11from numpy.typing import NDArray12 13from fastembed.common.types import NumpyArray14 15T = TypeVar("T")16 17 18def normalize(input_array: NumpyArray, p: int = 2, dim: int = 1, eps: float = 1e-12) -> NumpyArray:19 # Calculate the Lp norm along the specified dimension20 norm = np.linalg.norm(input_array, ord=p, axis=dim, keepdims=True)21 norm = np.maximum(norm, eps) # Avoid division by zero22 normalized_array = input_array / norm23 return normalized_array24 25 26def mean_pooling(input_array: NumpyArray, attention_mask: NDArray[np.int64]) -> NumpyArray:27 input_mask_expanded = np.expand_dims(attention_mask, axis=-1).astype(np.int64)28 input_mask_expanded = np.tile(input_mask_expanded, (1, 1, input_array.shape[-1]))29 sum_embeddings = np.sum(input_array * input_mask_expanded, axis=1)30 sum_mask = np.sum(input_mask_expanded, axis=1)31 pooled_embeddings = sum_embeddings / np.maximum(sum_mask, 1e-9)32 return pooled_embeddings33 34 35def iter_batch(iterable: Iterable[T], size: int) -> Iterable[list[T]]:36 """37 >>> list(iter_batch([1,2,3,4,5], 3))38 [[1, 2, 3], [4, 5]]39 """40 source_iter = iter(iterable)41 while source_iter:42 b = list(islice(source_iter, size))43 if len(b) == 0:44 break45 yield b46 47 48def define_cache_dir(cache_dir: str | None = None) -> Path:49 """50 Define the cache directory for fastembed51 """52 if cache_dir is None:53 default_cache_dir = os.path.join(tempfile.gettempdir(), "fastembed_cache")54 cache_path = Path(os.getenv("FASTEMBED_CACHE_PATH", default_cache_dir))55 else:56 cache_path = Path(cache_dir)57 cache_path.mkdir(parents=True, exist_ok=True)58 59 return cache_path60 61 62def get_all_punctuation() -> set[str]:63 return set(64 chr(i) for i in range(sys.maxunicode) if unicodedata.category(chr(i)).startswith("P")65 )66 67 68def remove_non_alphanumeric(text: str) -> str:69 return re.sub(r"[^\w\s]", " ", text, flags=re.UNICODE)70 