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使用
小结¶
提供数据在不同并行策略之间转换的标准接口。