diff --git a/haystack/nodes/retriever/_embedding_encoder.py b/haystack/nodes/retriever/_embedding_encoder.py index 775c813cbea..d12812ebd54 100644 --- a/haystack/nodes/retriever/_embedding_encoder.py +++ b/haystack/nodes/retriever/_embedding_encoder.py @@ -391,8 +391,6 @@ def save(self, save_dir: Union[Path, str]): class _OpenAIEmbeddingEncoder(_BaseEmbeddingEncoder): def __init__(self, retriever: "EmbeddingRetriever"): # See https://beta.openai.com/docs/guides/embeddings for more details - # OpenAI has a max seq length of 2048 tokens and unknown max batch size - self.max_seq_len = min(2048, retriever.max_seq_len) self.url = "https://api.openai.com/v1/embeddings" self.api_key = retriever.api_key self.batch_size = min(64, retriever.batch_size) @@ -400,11 +398,11 @@ def __init__(self, retriever: "EmbeddingRetriever"): model_class: str = next( (m for m in ["ada", "babbage", "davinci", "curie"] if m in retriever.embedding_model), "babbage" ) - self._setup_encoding_models(model_class, retriever.embedding_model) + self._setup_encoding_models(model_class, retriever.embedding_model, retriever.max_seq_len) self.tokenizer = AutoTokenizer.from_pretrained("gpt2") - def _setup_encoding_models(self, model_class: str, model_name: str): + def _setup_encoding_models(self, model_class: str, model_name: str, max_seq_len: int): """ Setup the encoding models for the retriever. """ @@ -412,9 +410,11 @@ def _setup_encoding_models(self, model_class: str, model_name: str): if "text-embedding" in model_name: self.query_encoder_model = model_name self.doc_encoder_model = model_name + self.max_seq_len = min(8191, max_seq_len) else: self.query_encoder_model = f"text-search-{model_class}-query-001" self.doc_encoder_model = f"text-search-{model_class}-doc-001" + self.max_seq_len = min(2046, max_seq_len) def _ensure_text_limit(self, text: str) -> str: """