-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmodel.py
More file actions
291 lines (255 loc) · 10.1 KB
/
Copy pathmodel.py
File metadata and controls
291 lines (255 loc) · 10.1 KB
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
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
import os
import base64
import configparser
import asyncio
from openai import OpenAI,AsyncOpenAI
import numpy as np
import cv2
def load_config():
"""加载配置文件"""
config = configparser.ConfigParser()
config.read('config.ini')
return config
def CLIENT(type):
"""根据类型初始化客户端"""
config = load_config()
if type in config:
client = OpenAI(
api_key=config[type]['api_key'],
base_url=config[type]['base_url'],
)
default_model = config[type]['default_model']
return client, default_model
else:
raise ValueError(f"Unsupported client type: {type}")
def async_CLIENT(type):
"""根据类型初始化异步客户端"""
config = load_config()
if type in config:
client = AsyncOpenAI(
api_key=config[type]['api_key'],
base_url=config[type]['base_url'],
)
default_model = config[type]['default_model']
return client, default_model
else:
raise ValueError(f"Unsupported client type: {type}")
def simple_llm(name: str, content: str, img_input = None, history = None, sys = None):
client, model = CLIENT(name)
# 处理图像输入(如果有)
img_url = None
if img_input is not None:
if isinstance(img_input, np.ndarray):
# 如果输入是numpy数组(OpenCV图像),将其转换为base64
success, buffer = cv2.imencode('.jpg', img_input)
if success:
img_b64_str = base64.b64encode(buffer).decode('utf-8')
img_url = f"data:image/jpeg;base64,{img_b64_str}"
else:
print("警告: 无法编码图像")
elif isinstance(img_input, str) and not img_input.startswith("http"):
# 处理本地文件路径
try:
with open(img_input, "rb") as image_file:
img_data = image_file.read()
img_b64_str = base64.b64encode(img_data).decode('utf-8')
img_url = f"data:image/jpeg;base64,{img_b64_str}"
except Exception as e:
print(f"读取图像文件时出错: {e}")
else:
img_url = img_input
# 准备消息列表
messages = []
# 添加系统提示(如果有)
if sys is not None:
messages.append({
"role": "system",
"content": sys
})
# 添加历史记录(如果有)
if history is not None:
messages.extend(history)
# 创建用户消息
if img_url:
# 多模态消息
user_message = {
"role": "user",
"content": [
{"type": "text", "text": content}
]
}
user_message["content"].append({
"type": "image_url",
"image_url": {"url": img_url}
})
else:
# 纯文本消息
user_message = {
"role": "user",
"content": content
}
messages.append(user_message)
# 调用API
try:
response = client.chat.completions.create(
model=model,
messages=messages,
)
return response.choices[0].message.content
except Exception as e:
print(f"Error: {e}")
return None
async def async_query(name, prompt, img_input=None, history=None, sys=None):
client, model = async_CLIENT(name)
messages = []
# 添加系统提示(如果有)
if sys is not None:
messages.append({
"role": "system",
"content": sys
})
# 添加历史记录(如果有)
if history is not None:
messages.extend(history)
# 处理图像输入(如果有)
img_url = None
if img_input is not None:
if isinstance(img_input, np.ndarray):
# 如果输入是numpy数组(OpenCV图像),将其转换为base64
success, buffer = cv2.imencode('.jpg', img_input)
if success:
img_b64_str = base64.b64encode(buffer).decode('utf-8')
img_url = f"data:image/jpeg;base64,{img_b64_str}"
else:
print("警告: 无法编码图像")
elif isinstance(img_input, str) and not img_input.startswith("http"):
# 处理本地文件路径
try:
with open(img_input, "rb") as image_file:
img_data = image_file.read()
img_b64_str = base64.b64encode(img_data).decode('utf-8')
img_url = f"data:image/jpeg;base64,{img_b64_str}"
except Exception as e:
print(f"读取图像文件时出错: {e}")
else:
img_url = img_input
# 创建用户消息
if img_url:
# 多模态消息
user_message = {
"role": "user",
"content": [
{"type": "text", "text": prompt}
]
}
user_message["content"].append({
"type": "image_url",
"image_url": {"url": img_url}
})
messages.append(user_message)
else:
# 纯文本消息
messages.append({"role": "user", "content": prompt})
# 调用API
try:
chat_completion = await client.chat.completions.create(
model=model,
messages=messages,
)
return chat_completion.choices[0].message.content
except Exception as e:
print(f"Async Error: {e}")
return None
async def async_simple_llm(name, queries, img_inputs=None, histories=None, sys=None):
tasks = []
# 处理图像输入
if img_inputs is not None and not isinstance(img_inputs, list):
img_inputs = [img_inputs] * len(queries)
elif img_inputs is not None and len(img_inputs) != len(queries):
if len(img_inputs) < len(queries):
img_inputs = img_inputs + [None] * (len(queries) - len(img_inputs))
else:
img_inputs = img_inputs[:len(queries)]
else:
img_inputs = [None] * len(queries) if img_inputs is None else img_inputs
# 处理历史记录
if histories is not None and not isinstance(histories, list):
histories = [histories] * len(queries)
elif histories is not None and len(histories) != len(queries):
if len(histories) < len(queries):
histories = histories + [None] * (len(queries) - len(histories))
else:
histories = histories[:len(queries)]
else:
histories = [None] * len(queries) if histories is None else histories
# 创建任务列表
tasks = [async_query(name, query, img, history, sys)
for query, img, history in zip(queries, img_inputs, histories)]
# 并行执行所有任务
results = await asyncio.gather(*tasks)
return results
class DialogueAgent:
def __init__(self, client_type,tools=None,system_msg=None):
self.client, self.default_model = CLIENT(client_type)
self.context = []
self.tools = tools
if system_msg is not None:
self.context.append({"role": "system", "content": system_msg})
def update_context(self, message, role="user",tool_call_id=None):
"""Update the dialogue context with a new message."""
if tool_call_id is not None and role == "tool":
self.context.append({"role": role, "content": message,"tool_call_id":tool_call_id})
else:
self.context.append({"role": role, "content": message})
def generate_response(self, user_input=None):
"""Generate a response based on the current context."""
# 更新上下文
if user_input is not None:
self.update_context(user_input, role="user")
# 发送请求给语言模型API
response = self.client.chat.completions.create(
model=self.default_model,
messages=self.context,
tools=self.tools
)
# 提取并更新上下文中的助手回复,内部更新直接更新context,不需要经过update_context这个接口
self.context.append(response.choices[0].message)
#如需得到对话内容,访问content字段
#如需得到工具调用,访问tool_calls字段
#如需得到推理内容,访问reasoning_content字段
return response.choices[0].message
def get_chat_content(self):
chat_content = []
for message in self.context:
entry = {}
role = message["role"] if isinstance(message, dict) else message.role
content = message["content"] if isinstance(message, dict) else message.content
# 处理assistant消息中可能包含的工具调用
if role == "assistant" and hasattr(message, "tool_calls") and message.tool_calls:
# 由于限定agent一次只会调用一个tool,所以只取第一个tool_call
tool_call = message.tool_calls[0]
entry = {"role": role,"content": f"{tool_call.function.name}:{tool_call.function.arguments}"}
# 处理常规消息
elif role == "user" or role == "system" or role == "assistant":
entry = {"role": role, "content": content}
# 处理工具调用消息
elif role == "tool":
entry = {"role": role, "content": content}
if entry: # 只添加非空条目
chat_content.append(entry)
return chat_content#得到一个列表,列表中每个元素是一个字典,字典中包含role和content字段
def change_model(self, new_model,keep_context=True,system_msg=None):
print(f"Changing model from {default_model} to {new_model}")
client, default_model = CLIENT(new_model)
self.client = client
self.default_model = default_model
if not keep_context:
self.context = []
if system_msg is not None:
self.context.append({"role": "system", "content": system_msg})
def clear_chat_content(self,clear_sys=True):
if clear_sys:
self.context = []
else:
#不清理系统提示词
self.context = [self.context[0]]