move TritonOps

This commit is contained in:
Daniel Han 2025-01-22 16:56:01 -08:00
commit 545e420ce6
2 changed files with 23 additions and 29 deletions

View file

@ -143,33 +143,10 @@ if Version(triton.__version__) >= Version("3.0.0"):
except: pass
else: from triton.common.build import libcuda_dirs
def fix_triton_ops():
# Check if triton.ops exists
try:
import triton.ops
except:
# Triton 3.2 removed triton.ops
from .matmul_perf_model import (
early_config_prune as _early_config_prune,
estimate_matmul_time as _estimate_matmul_time,
)
class PerfOps:
def __init__(self): return
@staticmethod
def early_config_prune(*args, **kwargs):
return _early_config_prune(*args, **kwargs)
@staticmethod
def estimate_matmul_time(*args, **kwargs):
return _estimate_matmul_time(*args, **kwargs)
pass
class TritonOps:
__slots__ = "matmul_perf_model",
def __init__(self): self.matmul_perf_model = PerfOps()
pass
triton.ops = TritonOps()
pass
pass
fix_triton_ops()
# Triton 3.2 removed triton.ops, so we shall fix it!
from .matmul_perf_model import TritonOps
try: import triton.ops
except: triton.ops = TritonOps()
# Try loading bitsandbytes and triton
import bitsandbytes as bnb
@ -210,7 +187,9 @@ except:
else: from triton.common.build import libcuda_dirs
cdequantize_blockwise_fp32 = bnb.functional.lib.cdequantize_blockwise_fp32
libcuda_dirs()
fix_triton_ops()
# Triton 3.2 removed triton.ops, so we shall fix it!
try: import triton.ops
except: triton.ops = TritonOps()
except:
warnings.warn(
"Unsloth: CUDA is not linked properly.\n"\

View file

@ -208,4 +208,19 @@ def early_config_prune(configs, named_args, **kwargs):
random_config = v[0][0]
random_config.num_stages = 2
pruned_configs.append(random_config)
return pruned_configs
return pruned_configs
class PerfOps:
def __init__(self): return
@staticmethod
def early_config_prune(*args, **kwargs):
return _early_config_prune(*args, **kwargs)
@staticmethod
def estimate_matmul_time(*args, **kwargs):
return _estimate_matmul_time(*args, **kwargs)
pass
class TritonOps:
__slots__ = "matmul_perf_model",
def __init__(self): self.matmul_perf_model = PerfOps()
pass