Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
mlx_pipeline.py257 linesDownload Raw Back to llms
1from __future__ import annotations2 3import logging4from typing import Any, Callable, Iterator, List, Mapping, Optional5 6from langchain_core.callbacks import CallbackManagerForLLMRun7from langchain_core.language_models.llms import LLM8from langchain_core.outputs import GenerationChunk9from pydantic import ConfigDict10 11DEFAULT_MODEL_ID = "mlx-community/quantized-gemma-2b"12 13logger = logging.getLogger(__name__)14 15 16class MLXPipeline(LLM):17    """MLX Pipeline API.18 19    To use, you should have the ``mlx-lm`` python package installed.20 21    Example using from_model_id:22        .. code-block:: python23 24            from langchain_community.llms import MLXPipeline25            pipe = MLXPipeline.from_model_id(26                model_id="mlx-community/quantized-gemma-2b",27                pipeline_kwargs={"max_tokens": 10, "temp": 0.7},28            )29    Example passing model and tokenizer in directly:30        .. code-block:: python31 32            from langchain_community.llms import MLXPipeline33            from mlx_lm import load34            model_id="mlx-community/quantized-gemma-2b"35            model, tokenizer = load(model_id)36            pipe = MLXPipeline(model=model, tokenizer=tokenizer)37    """38 39    model_id: str = DEFAULT_MODEL_ID40    """Model name to use."""41    model: Any = None  #: :meta private:42    """Model."""43    tokenizer: Any = None  #: :meta private:44    """Tokenizer."""45    tokenizer_config: Optional[dict] = None46    """47        Configuration parameters specifically for the tokenizer.48        Defaults to an empty dictionary.49    """50    adapter_file: Optional[str] = None51    """52        Path to the adapter file. If provided, applies LoRA layers to the model.53        Defaults to None.54    """55    lazy: bool = False56    """57        If False eval the model parameters to make sure they are58        loaded in memory before returning, otherwise they will be loaded59        when needed. Default: ``False``60    """61    pipeline_kwargs: Optional[dict] = None62    """63    Keyword arguments passed to the pipeline. Defaults include:64        - temp (float): Temperature for generation, default is 0.0.65        - max_tokens (int): Maximum tokens to generate, default is 100.66        - verbose (bool): Whether to output verbose logging, default is False.67        - formatter (Optional[Callable]): A callable to format the output.68          Default is None.69        - repetition_penalty (Optional[float]): The penalty factor for70          repeated sequences, default is None.71        - repetition_context_size (Optional[int]): Size of the context72          for applying repetition penalty, default is None.73        - top_p (float): The cumulative probability threshold for74          top-p filtering, default is 1.0.75 76    """77 78    model_config = ConfigDict(79        extra="forbid",80    )81 82    @classmethod83    def from_model_id(84        cls,85        model_id: str,86        tokenizer_config: Optional[dict] = None,87        adapter_file: Optional[str] = None,88        lazy: bool = False,89        pipeline_kwargs: Optional[dict] = None,90        **kwargs: Any,91    ) -> MLXPipeline:92        """Construct the pipeline object from model_id and task."""93        try:94            from mlx_lm import load95 96        except ImportError:97            raise ImportError(98                "Could not import mlx_lm python package. "99                "Please install it with `pip install mlx_lm`."100            )101 102        tokenizer_config = tokenizer_config or {}103        if adapter_file:104            model, tokenizer = load(105                model_id, tokenizer_config, adapter_path=adapter_file, lazy=lazy106            )107        else:108            model, tokenizer = load(model_id, tokenizer_config, lazy=lazy)109 110        _pipeline_kwargs = pipeline_kwargs or {}111        return cls(112            model_id=model_id,113            model=model,114            tokenizer=tokenizer,115            tokenizer_config=tokenizer_config,116            adapter_file=adapter_file,117            lazy=lazy,118            pipeline_kwargs=_pipeline_kwargs,119            **kwargs,120        )121 122    @property123    def _identifying_params(self) -> Mapping[str, Any]:124        """Get the identifying parameters."""125        return {126            "model_id": self.model_id,127            "tokenizer_config": self.tokenizer_config,128            "adapter_file": self.adapter_file,129            "lazy": self.lazy,130            "pipeline_kwargs": self.pipeline_kwargs,131        }132 133    @property134    def _llm_type(self) -> str:135        return "mlx_pipeline"136 137    def _call(138        self,139        prompt: str,140        stop: Optional[List[str]] = None,141        run_manager: Optional[CallbackManagerForLLMRun] = None,142        **kwargs: Any,143    ) -> str:144        try:145            from mlx_lm import generate146            from mlx_lm.sample_utils import make_logits_processors, make_sampler147 148        except ImportError:149            raise ImportError(150                "Could not import mlx_lm python package. "151                "Please install it with `pip install mlx_lm`."152            )153 154        pipeline_kwargs = kwargs.get("pipeline_kwargs", self.pipeline_kwargs) or {}155 156        temp: float = pipeline_kwargs.get("temp", 0.0)157        max_tokens: int = pipeline_kwargs.get("max_tokens", 100)158        verbose: bool = pipeline_kwargs.get("verbose", False)159        formatter: Optional[Callable] = pipeline_kwargs.get("formatter", None)160        repetition_penalty: Optional[float] = pipeline_kwargs.get(161            "repetition_penalty", None162        )163        repetition_context_size: Optional[int] = pipeline_kwargs.get(164            "repetition_context_size", None165        )166        top_p: float = pipeline_kwargs.get("top_p", 1.0)167        min_p: float = pipeline_kwargs.get("min_p", 0.0)168        min_tokens_to_keep: int = pipeline_kwargs.get("min_tokens_to_keep", 1)169 170        sampler = make_sampler(temp, top_p, min_p, min_tokens_to_keep)171        logits_processors = make_logits_processors(172            None, repetition_penalty, repetition_context_size173        )174 175        return generate(176            model=self.model,177            tokenizer=self.tokenizer,178            prompt=prompt,179            max_tokens=max_tokens,180            verbose=verbose,181            formatter=formatter,182            sampler=sampler,183            logits_processors=logits_processors,184        )185 186    def _stream(187        self,188        prompt: str,189        stop: Optional[List[str]] = None,190        run_manager: Optional[CallbackManagerForLLMRun] = None,191        **kwargs: Any,192    ) -> Iterator[GenerationChunk]:193        try:194            import mlx.core as mx195            from mlx_lm.sample_utils import make_logits_processors, make_sampler196            from mlx_lm.utils import generate_step197 198        except ImportError:199            raise ImportError(200                "Could not import mlx_lm python package. "201                "Please install it with `pip install mlx_lm`."202            )203 204        pipeline_kwargs = kwargs.get("pipeline_kwargs", self.pipeline_kwargs) or {}205 206        temp: float = pipeline_kwargs.get("temp", 0.0)207        max_new_tokens: int = pipeline_kwargs.get("max_tokens", 100)208        repetition_penalty: Optional[float] = pipeline_kwargs.get(209            "repetition_penalty", None210        )211        repetition_context_size: Optional[int] = pipeline_kwargs.get(212            "repetition_context_size", None213        )214        top_p: float = pipeline_kwargs.get("top_p", 1.0)215        min_p: float = pipeline_kwargs.get("min_p", 0.0)216        min_tokens_to_keep: int = pipeline_kwargs.get("min_tokens_to_keep", 1)217 218        prompt = self.tokenizer.encode(prompt, return_tensors="np")219 220        prompt_tokens = mx.array(prompt[0])221 222        eos_token_id = self.tokenizer.eos_token_id223        detokenizer = self.tokenizer.detokenizer224        detokenizer.reset()225 226        sampler = make_sampler(temp or 0.0, top_p, min_p, min_tokens_to_keep)227 228        logits_processors = make_logits_processors(229            None, repetition_penalty, repetition_context_size230        )231 232        for (token, prob), n in zip(233            generate_step(234                prompt=prompt_tokens,235                model=self.model,236                sampler=sampler,237                logits_processors=logits_processors,238            ),239            range(max_new_tokens),240        ):241            # identify text to yield242            text: Optional[str] = None243            detokenizer.add_token(token)244            detokenizer.finalize()245            text = detokenizer.last_segment246 247            # yield text, if any248            if text:249                chunk = GenerationChunk(text=text)250                if run_manager:251                    run_manager.on_llm_new_token(chunk.text)252                yield chunk253 254            # break if stop sequence found255            if token == eos_token_id or (stop is not None and text in stop):256                break257 
codekingpro/portable-devtools · Team Ai