跳转至

base.py — 分片管理器基类

文件概述

定义 BaseShardingManager,用于 Hybrid Engine 中数据在不同并行策略间的重分片。

核心类

BaseShardingManager

class BaseShardingManager:
    """数据分片管理器基类

    作为上下文管理器使用,在训练/推理切换时自动处理数据重分片。
    """

    def __enter__(self):
        """进入上下文:准备数据分片"""
        pass

    def __exit__(self, exc_type, exc_value, traceback):
        """退出上下文:恢复数据分片"""
        pass

    def preprocess_data(self, data: DataProto) -> DataProto:
        """预处理:将数据从一种并行布局转换为另一种"""
        return data

    def postprocess_data(self, data: DataProto) -> DataProto:
        """后处理:将结果数据恢复到原始并行布局"""
        return data

使用场景

with sharding_manager:
    # 数据已被重分片到当前并行策略
    data = sharding_manager.preprocess_data(data)
    result = model(data)
    result = sharding_manager.postprocess_data(result)
# 退出时自动恢复

与其他模块的关系

  • 被 fsdp_ulysses.py 继承实现
  • 被 fsdp_workers.py 使用

小结

提供数据在不同并行策略之间转换的标准接口。