diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index e03919f..6ee88b9 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -135,6 +135,8 @@ def _backend_supports( if isinstance(backend, FlashAttnBackend): if not flash_attn_available(): return False + if q.dtype not in (torch.float16, torch.bfloat16): + return False if q.size(1) == 1 and kv_cache is not None: return True return not (attn_mask is not None and not is_causal) diff --git a/tests/extension/test_backend.py b/tests/extension/test_backend.py index c056c83..61b3417 100644 --- a/tests/extension/test_backend.py +++ b/tests/extension/test_backend.py @@ -15,8 +15,8 @@ from astrai.extension import ( ) -def test_default_backend_is_torch_native(): - """Default is the highest-priority available backend (flash > cuda > torch).""" +def test_default_backend_resolves_to_available(): + """Default backend is the first available in cuda > flash > torch order.""" from astrai.extension.attention_backend import ( CudaBackend, FlashAttnBackend, @@ -26,7 +26,6 @@ def test_default_backend_is_torch_native(): backend = get_backend() assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend)) - assert isinstance(backend, type(_resolve_default_backend())) def test_attn_backend_context_with_enum(): diff --git a/tests/inference/test_scheduler.py b/tests/inference/test_scheduler.py index 0225fd4..8131b35 100644 --- a/tests/inference/test_scheduler.py +++ b/tests/inference/test_scheduler.py @@ -180,7 +180,7 @@ def test_scheduler_concurrent_get_stats(mock_model_and_tokenizer): def _make_real_scheduler(device): """Build a scheduler backed by a tiny real model for run_batch tests.""" cfg = make_rollout_config(max_position_embeddings=64) - model = AutoRegressiveLM(cfg).to(device=device).eval() + model = AutoRegressiveLM(cfg).to(device=device, dtype=torch.bfloat16).eval() tokenizer = FakeTokenizer() scheduler = InferenceScheduler( model=model,