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 定义了推理引擎与训练框架之间的标准接口:休眠/唤醒、权重更新、序列生成。