diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index d5cfdcfae..021457ffa 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -374,11 +374,6 @@ if os.environ.get("SDNQ_USE_TENSORWISE_FP8_MM", None) is None: else: use_tensorwise_fp8_matmul = os.environ.get("SDNQ_USE_TENSORWISE_FP8_MM", "0").lower() not in {"0", "false", "no"} -if os.environ.get("SDNQ_USE_CONTIGUOUS_MM", None) is None: - use_contiguous_mm = bool(use_openvino_mm or is_rdna2_and_older or devices.backend in {"ipex", "mps", "openvino", "zluda"}) -else: - use_contiguous_mm = bool(os.environ.get("SDNQ_USE_CONTIGUOUS_MM", "0").lower() not in {"0", "false", "no"}) - fp_mm_func = None int_mm_func = None @@ -412,6 +407,12 @@ if fp_mm_func is None: fp_mm_func = fp_mm_torch +if os.environ.get("SDNQ_USE_CONTIGUOUS_MM", None) is None: + use_contiguous_mm = bool(use_openvino_mm or is_rdna2_and_older or devices.backend in {"ipex", "mps", "openvino", "zluda"}) +else: + use_contiguous_mm = bool(os.environ.get("SDNQ_USE_CONTIGUOUS_MM", "0").lower() not in {"0", "false", "no"}) + + if use_torch_compile: torch._dynamo.config.recompile_limit = max(8192, getattr(torch._dynamo.config, "recompile_limit", 0)) torch._dynamo.config.cache_size_limit = max(8192, getattr(torch._dynamo.config, "cache_size_limit", 0))