Team Ai
Modelpublic

OneScience-Group/Medgemma

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes26downloads
model_runner.py281 linesDownload Raw Back to models
1# MedGemma vLLM 模型运行器2# 实现基于 vLLM 的推理引擎3 4import logging5from typing import Any, Dict, List, Optional, Set, Mapping6import numpy as np7 8logger = logging.getLogger(__name__)9 10 11class VLLMModelRunner:12    """13    基于 vLLM 的 MedGemma 模型运行器14    兼容 MedGemma serving_framework 的 ModelRunner 接口15    """16 17    def __init__(18        self,19        model_path: str,20        tokenizer_path: Optional[str] = None,21        gpu_memory_utilization: float = 0.9,22        max_model_len: Optional[int] = None,23        tensor_parallel_size: int = 1,24        trust_remote_code: bool = True,25    ):26        """27        初始化 vLLM 模型运行器28 29        Args:30            model_path: 模型权重路径31            tokenizer_path: Tokenizer 路径(默认与 model_path 相同)32            gpu_memory_utilization: GPU 内存使用率33            max_model_len: 最大序列长度34            tensor_parallel_size: Tensor 并行大小35            trust_remote_code: 是否信任远程代码36        """37        self.model_path = model_path38        self.tokenizer_path = tokenizer_path or model_path39 40        try:41            from vllm import LLM42            self.vllm_available = True43        except ImportError:44            logger.warning("vLLM not available. Install with: pip install vllm")45            self.vllm_available = False46            self.llm = None47            return48 49        logger.info(f"Initializing vLLM with model: {model_path}")50        self.llm = LLM(51            model=self.model_path,52            tokenizer=self.tokenizer_path,53            gpu_memory_utilization=gpu_memory_utilization,54            max_model_len=max_model_len,55            tensor_parallel_size=tensor_parallel_size,56            trust_remote_code=trust_remote_code,57        )58        logger.info("vLLM model loaded successfully")59 60    def run_model_multiple_output(61        self,62        model_input: Mapping[str, np.ndarray] | np.ndarray,63        model_name: str = "default",64        model_version: Optional[int] = None,65        model_output_keys: Optional[Set[str]] = None,66        parameters: Optional[Mapping[str, Any]] = None,67    ) -> Mapping[str, np.ndarray]:68        """69        运行推理(兼容 MedGemma ModelRunner 接口)70 71        Args:72            model_input: 输入数据(字典或 numpy 数组)73            model_name: 模型名称74            model_version: 模型版本75            model_output_keys: 输出键集合76            parameters: 推理参数77 78        Returns:79            输出字典80        """81        if not self.vllm_available or self.llm is None:82            raise RuntimeError("vLLM is not available")83 84        from vllm import SamplingParams85 86        parameters = parameters or {}87        model_output_keys = model_output_keys or {"text_output"}88 89        # 提取 prompt90        if isinstance(model_input, dict):91            prompt = model_input.get("prompt", "")92            if isinstance(prompt, np.ndarray):93                prompt = prompt.tobytes().decode('utf-8')94        else:95            # 如果是 numpy 数组,解码为字符串96            prompt = model_input.tobytes().decode('utf-8')97 98        # 创建采样参数99        sampling_params = SamplingParams(100            max_tokens=parameters.get("max_tokens", 500),101            temperature=parameters.get("temperature", 0.7),102            top_p=parameters.get("top_p", 0.9),103            top_k=parameters.get("top_k", -1),104            n=parameters.get("n", 1),105        )106 107        # 运行推理108        outputs = self.llm.generate([prompt], sampling_params)109 110        # 格式化输出111        result = {}112        if "text_output" in model_output_keys:113            result["text_output"] = np.array([114                output.text.encode('utf-8') for output in outputs[0].outputs115            ])116        if "num_input_tokens" in model_output_keys:117            result["num_input_tokens"] = np.array([len(outputs[0].prompt_token_ids)])118        if "num_output_tokens" in model_output_keys:119            result["num_output_tokens"] = np.array([120                len(output.token_ids) for output in outputs[0].outputs121            ])122 123        return result124 125    def generate(126        self,127        prompts: List[str],128        max_tokens: int = 500,129        temperature: float = 0.7,130        top_p: float = 0.9,131        top_k: int = -1,132        n: int = 1,133    ) -> List[Dict[str, Any]]:134        """135        简化的生成接口136 137        Args:138            prompts: 输入提示列表139            max_tokens: 最大生成 token 数140            temperature: 采样温度141            top_p: Nucleus 采样参数142            top_k: Top-K 采样参数143            n: 生成数量144 145        Returns:146            生成结果列表147        """148        if not self.vllm_available or self.llm is None:149            raise RuntimeError("vLLM is not available")150 151        from vllm import SamplingParams152 153        sampling_params = SamplingParams(154            max_tokens=max_tokens,155            temperature=temperature,156            top_p=top_p,157            top_k=top_k,158            n=n,159        )160 161        outputs = self.llm.generate(prompts, sampling_params)162 163        results = []164        for output in outputs:165            result = {166                "prompt": output.prompt,167                "outputs": [168                    {169                        "text": o.text,170                        "token_ids": o.token_ids,171                        "cumulative_logprob": o.cumulative_logprob,172                        "finish_reason": o.finish_reason,173                    }174                    for o in output.outputs175                ],176                "num_input_tokens": len(output.prompt_token_ids),177            }178            results.append(result)179 180        return results181 182 183class TransformersModelRunner:184    """185    基于 Transformers 的备用模型运行器186    当 vLLM 不可用时使用187    """188 189    def __init__(190        self,191        model_path: str,192        tokenizer_path: Optional[str] = None,193        device: str = "cuda",194        torch_dtype: str = "auto",195    ):196        """197        初始化 Transformers 模型运行器198 199        Args:200            model_path: 模型路径201            tokenizer_path: Tokenizer 路径202            device: 设备(cuda 或 cpu)203            torch_dtype: 数据类型204        """205        import torch206        from transformers import AutoModelForCausalLM, AutoTokenizer207 208        self.device = device209        self.tokenizer_path = tokenizer_path or model_path210 211        logger.info(f"Loading model with transformers: {model_path}")212        self.tokenizer = AutoTokenizer.from_pretrained(self.tokenizer_path)213        self.model = AutoModelForCausalLM.from_pretrained(214            model_path,215            torch_dtype=torch_dtype if torch_dtype != "auto" else "auto",216            device_map=device,217        )218        logger.info("Model loaded successfully")219 220    def generate(221        self,222        prompts: List[str],223        max_tokens: int = 500,224        temperature: float = 0.7,225        top_p: float = 0.9,226        top_k: int = 50,227        n: int = 1,228    ) -> List[Dict[str, Any]]:229        """230        生成文本231 232        Args:233            prompts: 输入提示列表234            max_tokens: 最大生成 token 数235            temperature: 采样温度236            top_p: Nucleus 采样参数237            top_k: Top-K 采样参数238            n: 生成数量239 240        Returns:241            生成结果列表242        """243        import torch244 245        results = []246        for prompt in prompts:247            inputs = self.tokenizer(prompt, return_tensors="pt").to(self.device)248 249            with torch.no_grad():250                outputs = self.model.generate(251                    **inputs,252                    max_new_tokens=max_tokens,253                    temperature=temperature,254                    top_p=top_p,255                    top_k=top_k,256                    num_return_sequences=n,257                    do_sample=temperature > 0,258                )259 260            generated_texts = [261                self.tokenizer.decode(output, skip_special_tokens=True)262                for output in outputs263            ]264 265            result = {266                "prompt": prompt,267                "outputs": [268                    {269                        "text": text,270                        "token_ids": None,271                        "cumulative_logprob": None,272                        "finish_reason": "stop",273                    }274                    for text in generated_texts275                ],276                "num_input_tokens": len(inputs["input_ids"][0]),277            }278            results.append(result)279 280        return results281