Files
AstrAI/tests/extension/conftest.py
T
ViperEkura 738cb8f128 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
2026-07-31 08:32:22 +08:00

31 lines
694 B
Python

"""Shared fixtures for extension tests."""
import pytest
import torch
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from tests.conftest import skip_no_kernel
D = 64
CFG = dict(
vocab_size=1000,
hidden_size=128,
num_attention_heads=2,
num_key_value_heads=1,
intermediate_size=256,
max_position_embeddings=64,
num_hidden_layers=2,
rms_norm_eps=1e-5,
attn_type="gqa",
ffn_type="mlp",
)
@pytest.fixture
def cuda_model():
config = AutoRegressiveLMConfig(**CFG)
model = AutoRegressiveLM(config).to(device="cuda", dtype=torch.bfloat16)
model.eval()
return model, config