PPTAgent 多智能体系统中Agent设计机理:数据结构,类设计与特点分析
·
代码地址:https://github.com/icip-cas/PPTAgent
一、整体设计架构
在PPTAgent系统中,Agent和AsyncAgent类是智能体的核心实现,负责与语言模型交互、管理对话历史、处理图像输入以及计算token成本。整体设计采用了面向对象的方式,将智能体的功能模块化,同时提供同步和异步两种接口,以适应不同的使用场景。
二、核心数据结构:Turn类
Turn类是表示对话回合的基本数据结构,用于记录智能体与用户之间的一次交互。
@dataclass
class Turn:
id: int
prompt: str
response: str
message: list
retry: int = -1
images: list[str] = None
input_tokens: int = 0
output_tokens: int = 0
embedding: Tensor = None
def to_dict(self):
return {k: v for k, v in asdict(self).items() if k != "embedding"}
def calc_token(self):
if self.images is not None:
self.input_tokens += calc_image_tokens(self.images)
self.input_tokens += len(ENCODING.encode(self.prompt))
self.output_tokens = len(ENCODING.encode(self.response))
def __eq__(self, other):
return self is other
设计特点:
- 数据封装:将对话回合的所有相关信息封装在一个类中
- token计算:提供
calc_token方法计算输入和输出的token数量 - 图像支持:支持存储图像路径并计算图像token
- 嵌入向量:存储对话回合的嵌入向量,用于相似性搜索
- 序列化:提供
to_dict方法便于序列化
三、Agent类设计分析
Agent类是智能体的核心实现,定义了智能体的基本行为和属性。
3.1 核心属性
class Agent:
def __init__(
self,
name: str,
llm_mapping: dict[str, LLM | AsyncLLM],
text_model: Optional[LLM | AsyncLLM] = None,
record_cost: bool = False,
config: Optional[dict] = None,
env: Optional[Environment] = None,
):
self.name = name
self.config = config
if self.config is None:
with open(package_join("roles", f"{name}.yaml"), encoding="utf-8") as f:
self.config = yaml.safe_load(f)
self.llm_mapping = llm_mapping
self.llm = self.llm_mapping[self.config["use_model"]]
self.model = self.llm.model
self.record_cost = record_cost
self.text_model = text_model
self.return_json = self.config.get("return_json", False)
self.system_message = self.config["system_prompt"]
self.prompt_args = set(self.config["jinja_args"])
self.env = env or Environment(undefined=StrictUndefined)
self.template = self.env.from_string(self.config["template"])
self.retry_template = Template("""...""")
self.input_tokens = 0
self.output_tokens = 0
self._history: list[Turn] = []
run_args = self.config.get("run_args", {})
self.llm.__call__ = partial(self.llm.__call__, **run_args)
self.system_tokens = len(ENCODING.encode(self.system_message))
3.2 核心方法
3.2.1 __call__方法:智能体调用入口
def __call__(
self,
images: list[str] = None,
recent: int = 0,
similar: int = 0,
**jinja_args,
):
if isinstance(images, str):
images = [images]
assert self.prompt_args == set(
jinja_args.keys()
), f"Invalid arguments, expected: {self.prompt_args}, got: {jinja_args.keys()}"
prompt = self.template.render(**jinja_args)
history = self.get_history(similar, recent, prompt)
history_msg = []
for turn in history:
history_msg.extend(turn.message)
response, message = self.llm(
prompt,
system_message=self.system_message,
history=history_msg,
images=images,
return_message=True,
)
turn = Turn(
id=self.next_turn_id,
prompt=prompt,
response=response,
message=message,
images=images,
)
return turn.id, self.__post_process__(response, history, turn, similar)
3.2.2 get_history方法:获取对话历史
def get_history(self, similar: int, recent: int, prompt: str):
history = self._history[-recent:] if recent > 0 else []
if similar > 0:
assert isinstance(self.text_model, LLM), "text_model must be a LLM"
embedding = self.text_model.get_embedding(prompt)
history.sort(key=lambda x: cosine_similarity(embedding, x.embedding))
for turn in history:
if len(history) > similar + recent:
break
if turn not in history:
history.append(turn)
history.sort(key=lambda x: x.id)
return history
3.2.3 retry方法:重试失败的回合
def retry(self, feedback: str, traceback: str, turn_id: int, error_idx: int):
assert error_idx > 0, "error_idx must be greater than 0"
prompt = self.retry_template.render(feedback=feedback, traceback=traceback)
history = [t for t in self._history if t.id == turn_id]
history_msg = []
for turn in history:
history_msg.extend(turn.message)
response, message = self.llm(
prompt,
history=history_msg,
return_message=True,
)
turn = Turn(
id=turn_id,
prompt=prompt,
response=response,
message=message,
retry=error_idx,
)
return self.__post_process__(response, history, turn)
3.2.4 __post_process__方法:处理响应
def __post_process__(
self, response: str, history: list[Turn], turn: Turn, similar: int = 0
) -> str | dict:
self._history.append(turn)
if similar > 0:
turn.embedding = self.text_model.get_embedding(turn.prompt)
if self.record_cost:
turn.calc_token()
self.calc_cost(history + [turn])
if self.return_json:
response = get_json_from_response(response)
return response
3.2.5 calc_cost方法:计算token成本
def calc_cost(self, turns: list[Turn]):
for turn in turns[:-1]:
self.input_tokens += turn.input_tokens
self.input_tokens += turn.output_tokens
self.input_tokens += turns[-1].input_tokens
self.output_tokens += turns[-1].output_tokens
self.input_tokens += self.system_tokens
3.3 其他重要方法
to_sync/to_async: 在同步和异步智能体间转换next_turn_id: 获取下一个回合的IDhistory: 获取排序后的对话历史__repr__: 提供智能体的字符串表示
四、AsyncAgent类设计分析
AsyncAgent类继承自Agent类,提供了异步接口实现。
4.1 核心属性
class AsyncAgent(Agent):
def __init__(
self,
name: str,
llm_mapping: dict[str, AsyncLLM],
text_model: Optional[AsyncLLM] = None,
record_cost: bool = False,
config: Optional[dict] = None,
env: Optional[Environment] = None,
):
super().__init__(name, llm_mapping, text_model, record_cost, config, env)
self.llm = self.llm.to_async()
4.2 核心异步方法
4.2.1 __call__方法:异步调用入口
async def __call__(
self,
images: list[str] = None,
recent: int = 0,
similar: int = 0,
**jinja_args,
):
if isinstance(images, str):
images = [images]
assert self.prompt_args == set(
jinja_args.keys()
), f"Invalid arguments, expected: {self.prompt_args}, got: {jinja_args.keys()}"
prompt = self.template.render(**jinja_args)
history = await self.get_history(similar, recent, prompt)
history_msg = []
for turn in history:
history_msg.extend(turn.message)
response, message = await self.llm(
prompt,
system_message=self.system_message,
history=history_msg,
images=images,
return_message=True,
)
turn = Turn(
id=self.next_turn_id,
prompt=prompt,
response=response,
message=message,
images=images,
)
return turn.id, await self.__post_process__(response, history, turn, similar)
4.2.2 get_history方法:异步获取对话历史
async def get_history(self, similar: int, recent: int, prompt: str):
history = self._history[-recent:] if recent > 0 else []
if similar > 0:
embedding = await self.text_model.get_embedding(prompt)
history.sort(key=lambda x: cosine_similarity(embedding, x.embedding))
for turn in history:
if len(history) > similar + recent:
break
if turn not in history:
history.append(turn)
history.sort(key=lambda x: x.id)
return history
4.2.3 retry方法:异步重试
async def retry(self, feedback: str, traceback: str, turn_id: int, error_idx: int):
assert error_idx > 0, "error_idx must be greater than 0"
prompt = self.retry_template.render(feedback=feedback, traceback=traceback)
history = [t for t in self._history if t.id == turn_id]
history_msg = []
for turn in history:
history_msg.extend(turn.message)
response, message = await self.llm(
prompt,
history=history_msg,
return_message=True,
)
turn = Turn(
id=turn_id,
prompt=prompt,
response=response,
message=message,
retry=error_idx,
)
return await self.__post_process__(response, history, turn)
4.2.4 __post_process__方法:异步处理响应
async def __post_process__(
self, response: str, history: list[Turn], turn: Turn, similar: int = 0
):
self._history.append(turn)
if similar > 0:
turn.embedding = await self.text_model.get_embedding(turn.prompt)
if self.record_cost:
turn.calc_token()
self.calc_cost(history + [turn])
if self.return_json:
response = get_json_from_response(response)
return response
五、设计特点分析
5.1 模块化设计
- 将智能体的功能拆分为多个模块:对话历史管理、LLM交互、token计算等
- 通过类继承实现同步和异步版本的代码复用
- 使用配置文件定义智能体的行为,提高灵活性
5.2 灵活性与可扩展性
- 支持多种LLM模型,通过
llm_mapping进行管理 - 支持图像输入,扩展了智能体的多模态能力
- 提供相似历史搜索功能,增强对话上下文理解
- 支持JSON格式输出,便于与其他系统集成
5.3 健壮性
- 提供重试机制,处理LLM响应失败的情况
- 严格验证输入参数,确保符合要求
- 详细记录对话历史和token成本,便于调试和优化
5.4 性能优化
- 通过异步接口提高并发处理能力
- 支持选择性加载对话历史,减少token消耗
- 图像token计算考虑了图像压缩,避免不必要的token消耗
六、总结
PPTAgent系统中的Agent和AsyncAgent类设计体现了模块化、灵活性和健壮性的原则。通过封装LLM交互、对话历史管理和token计算等功能,提供了简洁而强大的接口。异步实现进一步提高了系统的并发处理能力。
这种设计使得智能体能够适应不同的应用场景,处理复杂的幻灯片编辑任务,并与系统中的其他组件无缝集成。同时,系统还提供了丰富的错误处理和调试功能,便于开发和优化。
未来,通过进一步增强缓存机制、动态模型选择和成本控制等功能,可以使智能体的性能和灵活性得到更大的提升。
更多推荐
所有评论(0)