Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
model_laboratory.py99 linesDownload Raw Back to langchain_classic
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 
codekingpro/portable-devtools · Team Ai