torch.distributed.fsdp的limit_all_gathers参数作用?torch.distributed.fsdp的limit_all_gathers参数作用?
FSDP 在前向与反向时需要对当前层的参数做 all-gather,把各 rank 持有的分片拼成完整参数再计算。若多个层或多次 all-gather 同时进行,会并发占用大量显存(多份完整参数临时存在),容易 OOM;且 all-gather 是集体通信,过多并发也会增加调度与同步开销。
limit_all_gathers(或等价选项):限制同一时刻处于「已 all-gather、未释放」状态的参数量,即对并发 all-gather 做限流。实现上通常通过调度:只有当前「未释放的 all-gathered 参数」占用的显存低于某阈值或数量时,才允许下一层执行 all-gather;否则等待前面某层释放后再进行。这样用少量额外同步换显存峰值下降,避免因多段同时全量参数而 OOM,特别在层数多、参数大的模型上有效。
显存紧张或大模型时建议开启;会略微增加通信与调度序列化,但能显著提高可训模型规模或 batch 上限。具体参数名与默认值以当前 PyTorch FSDP 文档为准(如 limit_all_gathers=True 或带数值的配置)。
| 返回模块 | 返回总览 |