codekingpro/portable-devtools
114k
1"""Experiment with different models."""2 3from __future__ import annotations4 5from collections.abc import Sequence6 7from langchain_core.language_models.llms import BaseLLM8from langchain_core.prompts.prompt import PromptTemplate9from langchain_core.utils.input import get_color_mapping, print_text10 11from langchain_classic.chains.base import Chain12from langchain_classic.chains.llm import LLMChain13 14 15class ModelLaboratory:16 """A utility to experiment with and compare the performance of different models."""17 18 def __init__(self, chains: Sequence[Chain], names: list[str] | None = None):19 """Initialize the ModelLaboratory with chains to experiment with.20 21 Args:22 chains: A sequence of chains to experiment with.23 Each chain must have exactly one input and one output variable.24 names: Optional list of names corresponding to each chain.25 If provided, its length must match the number of chains.26 27 28 Raises:29 ValueError: If any chain is not an instance of `Chain`.30 ValueError: If a chain does not have exactly one input variable.31 ValueError: If a chain does not have exactly one output variable.32 ValueError: If the length of `names` does not match the number of chains.33 """34 for chain in chains:35 if not isinstance(chain, Chain):36 msg = ( # type: ignore[unreachable]37 "ModelLaboratory should now be initialized with Chains. "38 "If you want to initialize with LLMs, use the `from_llms` method "39 "instead (`ModelLaboratory.from_llms(...)`)"40 )41 raise ValueError(msg) # noqa: TRY00442 if len(chain.input_keys) != 1:43 msg = (44 "Currently only support chains with one input variable, "45 f"got {chain.input_keys}"46 )47 raise ValueError(msg)48 if len(chain.output_keys) != 1:49 msg = (50 "Currently only support chains with one output variable, "51 f"got {chain.output_keys}"52 )53 if names is not None and len(names) != len(chains):54 msg = "Length of chains does not match length of names."55 raise ValueError(msg)56 self.chains = chains57 chain_range = [str(i) for i in range(len(self.chains))]58 self.chain_colors = get_color_mapping(chain_range)59 self.names = names60 61 @classmethod62 def from_llms(63 cls,64 llms: list[BaseLLM],65 prompt: PromptTemplate | None = None,66 ) -> ModelLaboratory:67 """Initialize the ModelLaboratory with LLMs and an optional prompt.68 69 Args:70 llms: A list of LLMs to experiment with.71 prompt: An optional prompt to use with the LLMs.72 If provided, the prompt must contain exactly one input variable.73 74 Returns:75 An instance of `ModelLaboratory` initialized with LLMs.76 """77 if prompt is None:78 prompt = PromptTemplate(input_variables=["_input"], template="{_input}")79 chains = [LLMChain(llm=llm, prompt=prompt) for llm in llms]80 names = [str(llm) for llm in llms]81 return cls(chains, names=names)82 83 def compare(self, text: str) -> None:84 """Compare model outputs on an input text.85 86 If a prompt was provided with starting the laboratory, then this text will be87 fed into the prompt. If no prompt was provided, then the input text is the88 entire prompt.89 90 Args:91 text: input text to run all models on.92 """93 print(f"\033[1mInput:\033[0m\n{text}\n") # noqa: T20194 for i, chain in enumerate(self.chains):95 name = self.names[i] if self.names is not None else str(chain)96 print_text(name, end="\n")97 output = chain.run(text)98 print_text(output, color=self.chain_colors[str(i)], end="\n\n")99 