diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index c31417a..c053e69 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -492,7 +492,6 @@ class CudaBackend(AttentionBackend): kv_cache.k_buffer[layer_id].index_copy_(0, loc, k[:, 0]) kv_cache.v_buffer[layer_id].index_copy_(0, loc, v[:, 0]) - b = q.size(0) q_3d = q.squeeze(1) kv_indptr = kv_cache.kv_indptr diff --git a/tests/data/test_dataset.py b/tests/data/test_dataset.py index a6f1c2e..59300f9 100644 --- a/tests/data/test_dataset.py +++ b/tests/data/test_dataset.py @@ -6,7 +6,6 @@ import numpy as np import pytest import torch -from astrai.config.preprocess_config import PipelineConfig from astrai.dataset.dataset import ( DatasetFactory, GRPODataset, diff --git a/tests/extension/conftest.py b/tests/extension/conftest.py index 5f398fd..8e3e373 100644 --- a/tests/extension/conftest.py +++ b/tests/extension/conftest.py @@ -5,7 +5,6 @@ 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( diff --git a/tests/extension/test_backend.py b/tests/extension/test_backend.py index 1dabe7b..5db6f63 100644 --- a/tests/extension/test_backend.py +++ b/tests/extension/test_backend.py @@ -10,7 +10,6 @@ from astrai.extension import ( ATTN_BACKEND, AttentionBackendFactory, CudaBackend, - TorchNativeBackend, attn_backend, get_backend, ) diff --git a/tests/extension/test_backend_equivalence.py b/tests/extension/test_backend_equivalence.py index 9f9e59e..ef2c75a 100644 --- a/tests/extension/test_backend_equivalence.py +++ b/tests/extension/test_backend_equivalence.py @@ -31,7 +31,6 @@ def test_training_forward_matches_torch(cuda_model): or non-bf16 inputs it falls back to torch SDPA. Verify the fallback path matches the torch-native forward exactly. """ - import pytest model, _ = cuda_model input_ids = torch.randint(0, 1000, (2, 16), device="cuda") diff --git a/tests/parallel/test_broadcast_state_dict.py b/tests/parallel/test_broadcast_state_dict.py index 4d5a349..89a2f31 100644 --- a/tests/parallel/test_broadcast_state_dict.py +++ b/tests/parallel/test_broadcast_state_dict.py @@ -5,14 +5,12 @@ multi-rank environment without requiring multiple GPUs. """ import torch -import torch.distributed as dist -import torch.nn as nn from astrai.model.transformer import AutoRegressiveLM from astrai.parallel import get_rank, spawn_parallel_fn from astrai.parallel.executor import broadcast_state_dict, create_ref_model from astrai.trainer.strategy import GRPOStrategy -from tests.helpers import FakeExecutor, make_rollout_config +from tests.helpers import make_rollout_config def _broadcast_worker(): @@ -77,8 +75,6 @@ def _create_ref_model_worker(): ) assert ref is not None, f"rank {rank}: ref model is None" - # Every rank should have rank-0's weights, not its own - rank0_sd = model.state_dict() if rank == 0 else None # Broadcast rank-0's original weights for comparison if rank == 0: expected_sd = {k: v.clone() for k, v in model.state_dict().items()}