[v1] fix non-persistent buffer sync for init_on_rank0 (#10820)

This commit is contained in:
Hazeldxq
2026-09-08 16:00:35 +08:00
committed by GitHub
parent dced5f8804
commit 673048c6a5
2 changed files with 56 additions and 0 deletions

View File

@@ -373,6 +373,8 @@ class FSDP2Engine:
init_mode = getattr(model, "_init_mode", "init_on_default")
if init_mode == "init_on_rank0":
non_persistent_buffers = self._save_non_persistent_buffers(model) if self.rank == 0 else {}
if getattr(model.config, "tie_word_embeddings", False):
model.tie_weights()
@@ -391,6 +393,13 @@ class FSDP2Engine:
# Broadcast the full state dict from the global rank-0 process to all ranks in this group.
options = StateDictOptions(full_state_dict=True, cpu_offload=True, broadcast_from_rank0=True)
set_model_state_dict(model, full_sd, options=options)
self._restore_non_persistent_buffers(model, non_persistent_buffers)
if self.world_size > 1:
for module in model.modules():
for buffer_name in sorted(module._non_persistent_buffers_set):
buffer = getattr(module, buffer_name, None)
if buffer is not None:
torch.distributed.broadcast(buffer, src=0)
if self.rank == 0:
logger.info("init_on_rank0 sync complete.")

View File

@@ -18,12 +18,15 @@ Validates that the FSDP2 meta loading path behaves correctly for tied weights
and non-persistent buffers by comparing it with the standard non-meta path.
"""
from types import SimpleNamespace
import torch
from transformers import AutoConfig
from llamafactory.v1.accelerator.interface import DistributedInterface
from llamafactory.v1.config.arg_parser import get_args
from llamafactory.v1.core.model_engine import ModelEngine
from llamafactory.v1.plugins.trainer_plugins.distributed import fsdp2 as fsdp2_module
from llamafactory.v1.plugins.trainer_plugins.distributed.fsdp2 import FSDP2Engine
@@ -42,6 +45,50 @@ def collect_non_persistent_buffers(model):
return result
def test_fsdp2_rank0_syncs_non_persistent_buffers(monkeypatch):
expected = torch.tensor([1.0, 0.5, 0.25])
class Rank0Model(torch.nn.Module):
def __init__(self, rank):
super().__init__()
self.config = SimpleNamespace(tie_word_embeddings=False)
self._init_mode = "init_on_rank0"
self.weight = torch.nn.Parameter(torch.ones(1))
buffer = expected.clone() if rank == 0 else torch.empty_like(expected, device="meta")
self.register_buffer("inv_freq", buffer, persistent=False)
def to_empty(self, *, device, recurse=True):
self.inv_freq = torch.zeros_like(expected, device=device)
return self
def set_model_state_dict(model, state_dict, **kwargs):
assert torch.count_nonzero(model.inv_freq) == 0
monkeypatch.setattr(fsdp2_module, "get_current_accelerator", lambda: torch.device("cpu"))
monkeypatch.setattr(fsdp2_module, "set_model_state_dict", set_model_state_dict)
for rank in (0, 1):
engine = object.__new__(FSDP2Engine)
engine.rank = rank
engine.world_size = 2
monkeypatch.setattr(engine, "prepare_model", lambda model: model)
monkeypatch.setattr(engine, "_warmup_grad_norm", lambda model: None)
def broadcast(buffer, src):
assert src == 0
if rank == 0:
assert torch.equal(buffer, expected)
else:
assert torch.count_nonzero(buffer) == 0
buffer.copy_(expected)
monkeypatch.setattr(torch.distributed, "broadcast", broadcast)
model = engine.shard_model(Rank0Model(rank))
assert "inv_freq" not in model.state_dict()
assert torch.equal(model.inv_freq, expected)
def test_fsdp2_meta_loading_buffers_and_tied_weights():
"""Verify non-persistent buffers and tied weights consistency after meta load."""
# 1. Initialize DistributedInterface for single process