diff --git a/modules/intel/ipex/__init__.py b/modules/intel/ipex/__init__.py index 45138e491..0f62c0244 100644 --- a/modules/intel/ipex/__init__.py +++ b/modules/intel/ipex/__init__.py @@ -183,6 +183,7 @@ def ipex_init(): # pylint: disable=too-many-statements torch.cuda.is_fp16_supported = lambda *args, **kwargs: True torch.backends.cuda.is_built = lambda *args, **kwargs: True torch.version.cuda = "12.1" + torch.cuda.get_arch_list = lambda: ["ats-m150", "pvc"] torch.cuda.get_device_capability = lambda *args, **kwargs: [12,1] torch.cuda.get_device_properties.major = 12 torch.cuda.get_device_properties.minor = 1 diff --git a/modules/intel/ipex/hijacks.py b/modules/intel/ipex/hijacks.py index b76df2157..d0a9b80d9 100644 --- a/modules/intel/ipex/hijacks.py +++ b/modules/intel/ipex/hijacks.py @@ -96,14 +96,14 @@ def torch_bmm(input, mat2, *, out=None): return original_torch_bmm(input, mat2, out=out) @wraps(torch.nn.functional.scaled_dot_product_attention) -def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False): +def scaled_dot_product_attention(query, key, value, attn_mask=None, dropout_p=0.0, is_causal=False, **kwargs): if query.dtype != key.dtype: key = key.to(dtype=query.dtype) if query.dtype != value.dtype: value = value.to(dtype=query.dtype) if attn_mask is not None and query.dtype != attn_mask.dtype: attn_mask = attn_mask.to(dtype=query.dtype) - return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal) + return original_scaled_dot_product_attention(query, key, value, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, **kwargs) # A1111 FP16 original_functional_group_norm = torch.nn.functional.group_norm