fix: default backend race, raise on explicit fallback
- _default_backend lazy init protected with threading.Lock - attention() raises when explicit backend cannot handle call - FlashAttnBackend rejects prefill with non-None attn_mask - training test uses TORCH_NATIVE backend directly
This commit is contained in:
@@ -25,27 +25,26 @@ def _ws(pool: PagePool) -> InferenceWorkspace:
|
||||
|
||||
@skip_no_kernel
|
||||
def test_training_forward_matches_torch(cuda_model):
|
||||
"""Training forward (kv_cache=None) should produce identical logits.
|
||||
"""Training forward (kv_cache=None) uses torch-native SDPA.
|
||||
|
||||
CudaBackend is now safe as a default: for training (``kv_cache=None``)
|
||||
or non-bf16 inputs it falls back to torch SDPA. Verify the fallback
|
||||
path matches the torch-native forward exactly.
|
||||
CudaBackend does not support training (requires kv_cache).
|
||||
Torch-native backend must match default (which falls back to torch).
|
||||
"""
|
||||
|
||||
model, _ = cuda_model
|
||||
input_ids = torch.randint(0, 1000, (2, 16), device="cuda")
|
||||
|
||||
with torch.no_grad():
|
||||
out_torch = model(input_ids)
|
||||
out_default = model(input_ids)
|
||||
|
||||
with attn_backend(ATTN_BACKEND.CUDA):
|
||||
with attn_backend(ATTN_BACKEND.TORCH_NATIVE):
|
||||
with torch.no_grad():
|
||||
out_cuda = model(input_ids)
|
||||
out_torch = model(input_ids)
|
||||
|
||||
torch.testing.assert_close(
|
||||
out_cuda["logits"], out_torch["logits"], atol=1e-6, rtol=1e-6
|
||||
out_torch["logits"], out_default["logits"], atol=1e-6, rtol=1e-6
|
||||
)
|
||||
assert out_torch["logits"].shape[0] == 2
|
||||
assert out_default["logits"].shape[0] == 2
|
||||
|
||||
|
||||
@skip_no_kernel
|
||||
|
||||
Reference in New Issue
Block a user