跳转至

base.py — Rollout 抽象基类与注册表

文件概述

定义推理引擎的抽象基类 BaseRollout 和推理引擎注册表 _ROLLOUT_REGISTRY。

核心类

BaseRollout(ABC)

class BaseRollout(ABC):
    def __init__(self, config: RolloutConfig, model_config: HFModelConfig, device_mesh=None):
        self.config = config
        self.model_config = model_config
        self.device_mesh = device_mesh

    async def resume(self, tags: list[str]):
        """恢复 GPU 显存占用(wake up)"""
        pass

    async def update_weights(self, weights, global_steps=None, **kwargs):
        """更新推理引擎的模型权重"""
        pass

    async def release(self):
        """释放 GPU 显存(sleep)"""
        pass

    def generate_sequences(self, prompts: DataProto) -> DataProto:
        """同步生成序列(已不推荐)"""
        raise NotImplementedError

注册表

_ROLLOUT_REGISTRY = {}

def get_rollout_class(name: str):
    """根据名称获取 Rollout 类

    支持的名称:
    - "vllm" → vLLM ServerAdapter
    - "sglang" → SGLang ServerAdapter
    - "trtllm" → TensorRT-LLM ServerAdapter
    - "hf" → HuggingFace HFRollout
    """

与其他模块的关系

  • 被所有推理引擎适配器继承(vLLM, SGLang, TRT-LLM, HF)
  • resume/release 是 Hybrid Engine 切换的核心接口
  • update_weights 是训练-推理权重同步的核心接口

小结

BaseRollout 定义了推理引擎与训练框架之间的标准接口:休眠/唤醒、权重更新、序列生成。