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)