diff --git a/libs/langchain/langchain/chains/hyde/base.py b/libs/langchain/langchain/chains/hyde/base.py index 0d633246c..a9573f6b0 100644 --- a/libs/langchain/langchain/chains/hyde/base.py +++ b/libs/langchain/langchain/chains/hyde/base.py @@ -9,6 +9,7 @@ from typing import Any, Dict, List, Optional import numpy as np from langchain_core.embeddings import Embeddings from langchain_core.language_models import BaseLanguageModel +from langchain_core.prompts import BasePromptTemplate from langchain_core.pydantic_v1 import Extra from langchain.callbacks.manager import CallbackManagerForChainRun @@ -72,11 +73,21 @@ class HypotheticalDocumentEmbedder(Chain, Embeddings): cls, llm: BaseLanguageModel, base_embeddings: Embeddings, - prompt_key: str, + prompt_key: Optional[str] = None, + custom_prompt: Optional[BasePromptTemplate] = None, **kwargs: Any, ) -> HypotheticalDocumentEmbedder: - """Load and use LLMChain for a specific prompt key.""" - prompt = PROMPT_MAP[prompt_key] + """Load and use LLMChain with either a specific prompt key or custom prompt.""" + if custom_prompt is not None: + prompt = custom_prompt + elif prompt_key is not None and prompt_key in PROMPT_MAP: + prompt = PROMPT_MAP[prompt_key] + else: + raise ValueError( + f"Must specify prompt_key if custom_prompt not provided. Should be one " + f"of {list(PROMPT_MAP.keys())}." + ) + llm_chain = LLMChain(llm=llm, prompt=prompt) return cls(base_embeddings=base_embeddings, llm_chain=llm_chain, **kwargs)