代码地址: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

设计特点:

  1. 数据封装:将对话回合的所有相关信息封装在一个类中
  2. token计算:提供calc_token方法计算输入和输出的token数量
  3. 图像支持:支持存储图像路径并计算图像token
  4. 嵌入向量:存储对话回合的嵌入向量,用于相似性搜索
  5. 序列化:提供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: 获取下一个回合的ID
  • history: 获取排序后的对话历史
  • __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计算等功能,提供了简洁而强大的接口。异步实现进一步提高了系统的并发处理能力。

这种设计使得智能体能够适应不同的应用场景,处理复杂的幻灯片编辑任务,并与系统中的其他组件无缝集成。同时,系统还提供了丰富的错误处理和调试功能,便于开发和优化。

未来,通过进一步增强缓存机制、动态模型选择和成本控制等功能,可以使智能体的性能和灵活性得到更大的提升。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐