OneScience-Group/Medgemma
026
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 