codekingpro/portable-devtools
114k
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 