Merge branch 'main' into nightly

This commit is contained in:
Daniel Han 2025-07-22 03:34:55 -07:00
commit 1476113115

View file

@ -86,11 +86,13 @@ from triton import __version__ as triton_version
HAS_XFORMERS = xformers is not None
BlockDiagonalCausalMask = xformers.attn_bias.BlockDiagonalCausalMask if HAS_XFORMERS else None
def clean_gpu_cache():
if DEVICE_TYPE == "xpu":
torch.xpu.empty_cache()
else:
torch.cuda.empty_cache()
if DEVICE_TYPE == "xpu":
clean_gpu_cache = torch.xpu.empty_cache
get_current_device = torch.xpu.current_device
else:
clean_gpu_cache = torch.cuda.empty_cache
get_current_device = torch.cuda.current_device
pass
def original_apply_qkv(self, X):
Q = self.q_proj(X)
@ -342,9 +344,9 @@ def LlamaAttention_fast_forward_inference(
A = torch_matmul(A, Vnn, out = Qn)
else:
if SDPA_HAS_GQA:
A = scaled_dot_product_attention(Qn, Knn, Vnn, attn_mask = attention_mask, is_causal = False, enable_gqa = True)
A = scaled_dot_product_attention(Qn, Knn, Vnn, attn_mask = attention_mask, is_causal = is_causal, enable_gqa = True)
else:
A = scaled_dot_product_attention(Qn, Knn, Vnn, attn_mask = attention_mask, is_causal = False)
A = scaled_dot_product_attention(Qn, Knn, Vnn, attn_mask = attention_mask, is_causal = is_causal)
pass
A = A.transpose(1, 2)
A = A.reshape(bsz, 1, attention_size)
@ -1365,8 +1367,8 @@ class LlamaRotaryEmbedding(torch.nn.Module):
self._set_cos_sin_cache(seq_len=self.current_rope_size, device=torch.device(device_idx), dtype=torch.get_default_dtype())
# dummy so that patch_utils doesn't fail for now
self.cos_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.sin_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.cos_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
self.sin_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
pass
def _set_cos_sin_cache(self, seq_len, device, dtype):
@ -1402,7 +1404,7 @@ class LlamaRotaryEmbedding(torch.nn.Module):
def get_cached(self, seq_len = None, device_index = None):
if device_index is None:
device_index = torch.cuda.current_device()
device_index = get_current_device()
return self.multi_gpu_cos_cached[device_index], self.multi_gpu_sin_cached[device_index]
pass
@ -1484,8 +1486,8 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
self._set_cos_sin_cache(seq_len=self.current_rope_size, device=torch.device(device_idx), dtype=torch.get_default_dtype())
# dummy so that patch_utils doesn't fail for now
self.cos_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.sin_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.cos_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
self.sin_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
pass
def _set_cos_sin_cache(self, seq_len, device, dtype):
@ -1518,7 +1520,7 @@ class LlamaExtendedRotaryEmbedding(torch.nn.Module):
def get_cached(self, seq_len = None, device_index = None):
if device_index is None:
device_index = torch.cuda.current_device()
device_index = get_current_device()
return self.multi_gpu_cos_cached[device_index], self.multi_gpu_sin_cached[device_index]
pass
@ -1631,10 +1633,10 @@ class LongRopeRotaryEmbedding(torch.nn.Module):
self.multi_gpu_short_sin_cached[device_idx] = sin_cached
# dummy so that patch_utils doesn't fail for now
self.short_cos_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.short_sin_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.long_cos_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.long_sin_cached = torch.empty(1, device=torch.cuda.current_device(), dtype=torch.get_default_dtype())
self.short_cos_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
self.short_sin_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
self.long_cos_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
self.long_sin_cached = torch.empty(1, device=get_current_device(), dtype=torch.get_default_dtype())
pass
def _set_cos_sin_cache(self, seq_len, device, dtype):
@ -1675,7 +1677,7 @@ class LongRopeRotaryEmbedding(torch.nn.Module):
def get_cached(self, seq_len = None, device_index = None):
if device_index is None:
device_index = torch.cuda.current_device()
device_index = get_current_device()
if seq_len is not None and seq_len < self.original_max_position_embeddings:
return self.multi_gpu_short_cos_cached[device_index], self.multi_gpu_short_sin_cached[device_index]
return self.multi_gpu_long_cos_cached[device_index], self.multi_gpu_long_sin_cached[device_index]