fix: broadcast ref/old model state_dict for FSDP

- Add broadcast_state_dict to sync state_dict from rank-0 to all ranks
- Fix create_ref_model returning None on non-rank-0 under FSDP
- Fix sync_old_model only updating old_model on rank-0 under FSDP
- Split skip_no_cuda/skip_no_kernel markers and hoist to top-level conftest
- Add distributed tests for broadcast_state_dict and create_ref_model
This commit is contained in:
2026-07-31 08:32:22 +08:00
parent 28d1bd07cf
commit 738cb8f128
9 changed files with 248 additions and 17 deletions
+6
View File
@@ -7,10 +7,16 @@ import pytest
import torch
from tokenizers import Tokenizer, models, pre_tokenizers, trainers
from astrai.extension import KERNEL_NAMES, is_available
from astrai.model.transformer import AutoRegressiveLM
from astrai.tokenize import AutoTokenizer
from tests.helpers import TINY_CONFIG, RandomTokenDataset, make_tiny_config
CUDA_AVAIL = torch.cuda.is_available()
KERNEL_AVAIL = CUDA_AVAIL and all(is_available(k) for k in KERNEL_NAMES)
skip_no_cuda = pytest.mark.skipif(not CUDA_AVAIL, reason="CUDA not available")
skip_no_kernel = pytest.mark.skipif(not KERNEL_AVAIL, reason="CUDA kernels not built")
def pytest_configure(config):
config.addinivalue_line("markers", "slow: marks tests as slow")