fix: skip CUDA rotary kernel when grad is enabled
- apply_rotary_emb now checks torch.is_grad_enabled() before dispatching to CUDA kernel - Training (grad enabled) uses torch complex multiply path which supports autograd backward - Inference (inference_mode/no_grad) uses CUDA kernel as before - Without this fix, training backward would crash — the CUDA kernel has no autograd backward()
This commit is contained in:
@@ -43,6 +43,11 @@ def apply_rotary_emb(x: Tensor, rotary_emb: tuple[Tensor, Tensor]) -> Tensor:
|
||||
[batch, seq_len, n_heads, head_dim] (bf16)
|
||||
"""
|
||||
cos, sin = rotary_emb
|
||||
if _cuda_available() and x.is_cuda and x.dtype == torch.bfloat16:
|
||||
if (
|
||||
_cuda_available()
|
||||
and not torch.is_grad_enabled()
|
||||
and x.is_cuda
|
||||
and x.dtype == torch.bfloat16
|
||||
):
|
||||
return _cuda_rotary(x, cos, sin)
|
||||
return _torch_apply(x, cos, sin)
|
||||
|
||||
Reference in New Issue
Block a user