File size: 6,283 Bytes
fedd8d3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 | # MedGemma 推理包装器
# 包装 MedGemma 原始 predictor.py 的推理逻辑
import logging
import os
import sys
from typing import Any, Dict, List, Optional
logger = logging.getLogger(__name__)
class MedGemmaPredictor:
"""
MedGemma 推理包装器
包装原始 MedGemma predictor 逻辑,提供 OneScience 兼容接口
"""
def __init__(self, model_runner: Any, configs: Any):
"""
初始化推理包装器
Args:
model_runner: 模型运行器(VLLMModelRunner 或 TransformersModelRunner)
configs: 配置对象
"""
self.model_runner = model_runner
self.configs = configs
# 尝试导入原始 MedGemma 组件(如果可用)
self._init_medgemma_components()
def _init_medgemma_components(self):
"""初始化 MedGemma 原始组件"""
try:
# 添加 MedGemma 原始代码路径到 sys.path
medgemma_base = os.path.abspath(
os.path.join(os.path.dirname(__file__), "..", "..", "..", "..", "..", "..", "medgemma", "python")
)
if os.path.exists(medgemma_base) and medgemma_base not in sys.path:
sys.path.insert(0, medgemma_base)
logger.info(f"Added MedGemma path: {medgemma_base}")
# 尝试导入 MedGemma predictor 组件
try:
from serving import predictor
self.has_original_predictor = True
logger.info("Successfully imported original MedGemma predictor")
except ImportError as e:
logger.warning(f"Could not import original MedGemma predictor: {e}")
self.has_original_predictor = False
except Exception as e:
logger.warning(f"Error initializing MedGemma components: {e}")
self.has_original_predictor = False
def predict(
self,
messages: List[Dict[str, Any]],
max_tokens: int = 500,
temperature: float = 0.7,
top_p: float = 0.9,
n: int = 1,
) -> Dict[str, Any]:
"""
运行推理
Args:
messages: OpenAI Chat Completion 格式的消息列表
max_tokens: 最大生成 token 数
temperature: 采样温度
top_p: Nucleus 采样参数
n: 生成数量
Returns:
OpenAI 兼容格式的响应
"""
# 转换消息为 prompt
prompt = self._messages_to_prompt(messages)
# 运行模型推理
results = self.model_runner.generate(
prompts=[prompt],
max_tokens=max_tokens,
temperature=temperature,
top_p=top_p,
n=n,
)
# 格式化响应为 OpenAI 格式
return self._format_openai_response(results[0], messages)
def _messages_to_prompt(self, messages: List[Dict[str, Any]]) -> str:
"""
将 OpenAI 消息格式转换为 prompt
Args:
messages: 消息列表
Returns:
格式化的 prompt 字符串
"""
prompt_parts = []
for message in messages:
role = message.get("role", "user")
content = message.get("content", "")
# 处理不同角色的消息
if role == "system":
prompt_parts.append(f"System: {content}")
elif role == "user":
prompt_parts.append(f"User: {content}")
elif role == "assistant":
prompt_parts.append(f"Assistant: {content}")
else:
prompt_parts.append(f"{role}: {content}")
# 添加 Assistant 前缀以开始生成
prompt_parts.append("Assistant:")
return "\n".join(prompt_parts)
def _format_openai_response(
self,
result: Dict[str, Any],
messages: List[Dict[str, Any]]
) -> Dict[str, Any]:
"""
将模型输出格式化为 OpenAI Chat Completion 格式
Args:
result: 模型生成结果
messages: 原始消息
Returns:
OpenAI 格式的响应
"""
import time
import uuid
choices = []
for idx, output in enumerate(result["outputs"]):
choice = {
"index": idx,
"message": {
"role": "assistant",
"content": output["text"].replace(result["prompt"], "").strip(),
},
"finish_reason": output.get("finish_reason", "stop"),
}
choices.append(choice)
response = {
"id": f"chatcmpl-{uuid.uuid4().hex[:8]}",
"object": "chat.completion",
"created": int(time.time()),
"model": self.configs.model.variant,
"choices": choices,
"usage": {
"prompt_tokens": result.get("num_input_tokens", 0),
"completion_tokens": sum(
len(output.get("token_ids", [])) if output.get("token_ids") else 0
for output in result["outputs"]
),
"total_tokens": result.get("num_input_tokens", 0) + sum(
len(output.get("token_ids", [])) if output.get("token_ids") else 0
for output in result["outputs"]
),
},
}
return response
def predict_with_images(
self,
messages: List[Dict[str, Any]],
images: List[Any],
max_tokens: int = 500,
temperature: float = 0.7,
top_p: float = 0.9,
) -> Dict[str, Any]:
"""
多模态推理(文本 + 图像)
Args:
messages: 消息列表
images: 图像列表
max_tokens: 最大生成 token 数
temperature: 采样温度
top_p: Nucleus 采样参数
Returns:
响应字典
"""
# TODO: 实现多模态推理
# 这需要集成 MedGemma 的图像处理逻辑
logger.warning("Multimodal inference not yet fully implemented")
# 暂时只处理文本
return self.predict(messages, max_tokens, temperature, top_p)
|