From b35cb28eeb7bdbed0d3ddbb823c9c9affe8e324d Mon Sep 17 00:00:00 2001 From: NirantK Date: Mon, 16 Oct 2023 17:39:21 +0530 Subject: [PATCH] * refactor(embedding.py): reorder import statements in alphabetical order * feat(embedding.py): add optional 'threads' parameter to DefaultEmbedding constructor --- fastembed/embedding.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/fastembed/embedding.py b/fastembed/embedding.py index 55f3ff4..7e99283 100644 --- a/fastembed/embedding.py +++ b/fastembed/embedding.py @@ -6,7 +6,7 @@ from abc import ABC, abstractmethod from itertools import islice from multiprocessing import get_all_start_methods from pathlib import Path -from typing import Dict, Iterable, List, Union, Generator, Any, Tuple +from typing import Any, Dict, Generator, Iterable, List, Optional, Tuple, Union import numpy as np import onnxruntime as ort @@ -14,7 +14,7 @@ import requests from tokenizers import AddedToken, Tokenizer from tqdm import tqdm -from fastembed.parallel_processor import Worker, ParallelWorkerPool +from fastembed.parallel_processor import ParallelWorkerPool, Worker def iter_batch(iterable: Union[Iterable, Generator], size: int) -> Iterable: @@ -471,7 +471,8 @@ class DefaultEmbedding(FlagEmbedding): self, model_name: str = "BAAI/bge-small-en-v1.5", max_length: int = 512, - cache_dir: str = None, + cache_dir: Optional[str] = None, + threads: Optional[int] = None, ): super().__init__(model_name, max_length=max_length, cache_dir=cache_dir, threads=threads)