diff --git a/src/candle/_cython/_functional_ops.pyx b/src/candle/_cython/_functional_ops.pyx index 536fdba0..6f0eef38 100644 --- a/src/candle/_cython/_functional_ops.pyx +++ b/src/candle/_cython/_functional_ops.pyx @@ -16,6 +16,12 @@ cdef object _py_mul_fn = None cdef object _py_matmul_fn = None cdef object _py_sub_fn = None cdef object _py_div_fn = None +cdef object _py_atan2_fn = None +cdef object _py_pow_fn = None +cdef object _py_remainder_fn = None +cdef object _py_fmod_fn = None +cdef object _py_logaddexp_fn = None +cdef object _py_logaddexp2_fn = None cdef object _py_relu_fn = None cdef object _py_neg_fn = None @@ -24,8 +30,15 @@ cdef object _npu_add_fn = None cdef object _npu_mul_fn = None cdef object _npu_sub_fn = None cdef object _npu_div_fn = None +cdef object _npu_atan2_fn = None +cdef object _npu_pow_tt_fn = None +cdef object _npu_remainder_fn = None +cdef object _npu_fmod_fn = None +cdef object _npu_logaddexp_fn = None +cdef object _npu_logaddexp2_fn = None cdef object _grad_mode_state = None cdef object _is_functionalize_fn = None +cdef object _use_soc_fallback_fn = None cdef object _current_pipeline_fn = None cdef bint _npu_refs_loaded = False @@ -82,25 +95,40 @@ cdef void _collect_types(object val, object types): cdef inline void _ensure_originals(): - global _py_add_fn, _py_mul_fn, _py_matmul_fn, _py_sub_fn, _py_div_fn - global _py_relu_fn, _py_neg_fn + global _py_add_fn, _py_mul_fn, _py_matmul_fn, _py_sub_fn, _py_div_fn, _py_atan2_fn + global _py_pow_fn, _py_remainder_fn, _py_fmod_fn, _py_logaddexp_fn, _py_logaddexp2_fn + global _py_relu_fn, _py_neg_fn, _use_soc_fallback_fn _ensure_dispatch() if _py_add_fn is None: - from candle._functional import _py_add, _py_mul, _py_matmul, _py_sub, _py_div, _py_relu, _py_neg + from candle._functional import ( + _py_add, _py_mul, _py_matmul, _py_sub, _py_div, _py_atan2, + _py_pow, _py_remainder, _py_fmod, _py_logaddexp, _py_logaddexp2, + _py_relu, _py_neg, + ) + from candle._backends.npu.ops._helpers import _use_soc_fallback as _usf _py_add_fn = _py_add _py_mul_fn = _py_mul _py_matmul_fn = _py_matmul _py_sub_fn = _py_sub _py_div_fn = _py_div + _py_atan2_fn = _py_atan2 + _py_pow_fn = _py_pow + _py_remainder_fn = _py_remainder + _py_fmod_fn = _py_fmod + _py_logaddexp_fn = _py_logaddexp + _py_logaddexp2_fn = _py_logaddexp2 _py_relu_fn = _py_relu _py_neg_fn = _py_neg + _use_soc_fallback_fn = _usf cdef inline void _ensure_npu_refs(): """Load NPU op refs and guard state once.""" - global _npu_add_fn, _npu_mul_fn, _npu_sub_fn, _npu_div_fn + global _npu_add_fn, _npu_mul_fn, _npu_sub_fn, _npu_div_fn, _npu_atan2_fn + global _npu_pow_tt_fn, _npu_remainder_fn, _npu_fmod_fn + global _npu_logaddexp_fn, _npu_logaddexp2_fn global _grad_mode_state, _is_functionalize_fn, _current_pipeline_fn global _npu_refs_loaded @@ -120,6 +148,30 @@ cdef inline void _ensure_npu_refs(): from candle._cython._npu_ops import fast_div as _ndiv # pylint: disable=import-error,no-name-in-module except ImportError: from candle._backends.npu.ops import div as _ndiv + try: + from candle._cython._npu_ops import fast_atan2 as _natan2 # pylint: disable=import-error,no-name-in-module + except ImportError: + from candle._backends.npu.ops import atan2 as _natan2 + try: + from candle._cython._npu_ops import fast_pow_tensor_tensor as _npow # pylint: disable=import-error,no-name-in-module + except ImportError: + _npow = None + try: + from candle._cython._npu_ops import fast_remainder as _nrem # pylint: disable=import-error,no-name-in-module + except ImportError: + _nrem = None + try: + from candle._cython._npu_ops import fast_fmod as _nfmod # pylint: disable=import-error,no-name-in-module + except ImportError: + _nfmod = None + try: + from candle._cython._npu_ops import fast_logaddexp as _nlae # pylint: disable=import-error,no-name-in-module + except ImportError: + _nlae = None + try: + from candle._cython._npu_ops import fast_logaddexp2 as _nlae2 # pylint: disable=import-error,no-name-in-module + except ImportError: + _nlae2 = None from candle.autograd.grad_mode import _GRAD_MODE_STATE as _gms from candle._dispatch.functionalize import is_functionalize_enabled as _ife from candle._dispatch.pipeline import current_pipeline as _cp @@ -128,6 +180,12 @@ cdef inline void _ensure_npu_refs(): _npu_mul_fn = _nmul _npu_sub_fn = _nsub _npu_div_fn = _ndiv + _npu_atan2_fn = _natan2 + _npu_pow_tt_fn = _npow + _npu_remainder_fn = _nrem + _npu_fmod_fn = _nfmod + _npu_logaddexp_fn = _nlae + _npu_logaddexp2_fn = _nlae2 _grad_mode_state = _gms _is_functionalize_fn = _ife _current_pipeline_fn = _cp @@ -318,6 +376,140 @@ def div(a, b, *, rounding_mode=None): return _dispatch_fn("true_divide", None, a, b) +def atan2(a, b): + """Fast atan2: skip __torch_function__ when both args are base Tensor.""" + cdef object r + + _ensure_originals() + + if _is_base_tensor(a) and (_is_base_tensor(b) or not hasattr(b, "__torch_function__")): + if hasattr(b, "shape") and _is_npu_tensor_pair(a, b): + if _use_soc_fallback_fn is not None and _use_soc_fallback_fn("atan2"): + return _dispatch_fn("atan2", None, a, b) + _ensure_npu_refs() + if _npu_fast_ok(a, b): + return _npu_atan2_fn(a, b) + return _dispatch_fn("atan2", None, a, b) + + r = _handle_torch_function(_py_atan2_fn, (a, b), {}) + if r is not NotImplemented: + return r + + return _dispatch_fn("atan2", None, a, b) + + + +def pow(a, b): + """Fast pow: skip __torch_function__ when both args are base Tensor.""" + cdef object r + + _ensure_originals() + + if _is_base_tensor(a) and (_is_base_tensor(b) or not hasattr(b, "__torch_function__")): + # Only tensor-tensor path gets the fast Cython path + if hasattr(b, "shape") and _is_npu_tensor_pair(a, b): + _ensure_npu_refs() + if _npu_fast_ok(a, b) and _npu_pow_tt_fn is not None: + return _npu_pow_tt_fn(a, b) + return _dispatch_fn("pow", None, a, b) + + r = _handle_torch_function(_py_pow_fn, (a, b), {}) + if r is not NotImplemented: + return r + + return _dispatch_fn("pow", None, a, b) + + +def remainder(a, b): + """Fast remainder: skip __torch_function__ when both args are base Tensor. + + Falls back to dispatch for SoC-specific fallback or scalar inputs. + """ + cdef object r + + _ensure_originals() + + if _is_base_tensor(a) and (_is_base_tensor(b) or not hasattr(b, "__torch_function__")): + # remainder has SoC fallback and scalar normalization in the backend; + # only take the Cython fast path for tensor-tensor on non-fallback SoCs. + if hasattr(b, "shape") and _is_npu_tensor_pair(a, b): + if _use_soc_fallback_fn is not None and _use_soc_fallback_fn("remainder"): + return _dispatch_fn("remainder", None, a, b) + _ensure_npu_refs() + if _npu_fast_ok(a, b) and _npu_remainder_fn is not None: + return _npu_remainder_fn(a, b) + return _dispatch_fn("remainder", None, a, b) + + r = _handle_torch_function(_py_remainder_fn, (a, b), {}) + if r is not NotImplemented: + return r + + return _dispatch_fn("remainder", None, a, b) + + +def fmod(a, b): + """Fast fmod: skip __torch_function__ when both args are base Tensor. + + Falls back to dispatch for scalar inputs. + """ + cdef object r + + _ensure_originals() + + if _is_base_tensor(a) and (_is_base_tensor(b) or not hasattr(b, "__torch_function__")): + if hasattr(b, "shape") and _is_npu_tensor_pair(a, b): + _ensure_npu_refs() + if _npu_fast_ok(a, b) and _npu_fmod_fn is not None: + return _npu_fmod_fn(a, b) + return _dispatch_fn("fmod", None, a, b) + + r = _handle_torch_function(_py_fmod_fn, (a, b), {}) + if r is not NotImplemented: + return r + + return _dispatch_fn("fmod", None, a, b) + + +def logaddexp(a, b): + """Fast logaddexp: skip __torch_function__ when both args are base Tensor.""" + cdef object r + + _ensure_originals() + + if _is_base_tensor(a) and (_is_base_tensor(b) or not hasattr(b, "__torch_function__")): + if _is_npu_tensor_pair(a, b): + _ensure_npu_refs() + if _npu_fast_ok(a, b) and _npu_logaddexp_fn is not None: + return _npu_logaddexp_fn(a, b) + return _dispatch_fn("logaddexp", None, a, b) + + r = _handle_torch_function(_py_logaddexp_fn, (a, b), {}) + if r is not NotImplemented: + return r + + return _dispatch_fn("logaddexp", None, a, b) + + +def logaddexp2(a, b): + """Fast logaddexp2: skip __torch_function__ when both args are base Tensor.""" + cdef object r + + _ensure_originals() + + if _is_base_tensor(a) and (_is_base_tensor(b) or not hasattr(b, "__torch_function__")): + if _is_npu_tensor_pair(a, b): + _ensure_npu_refs() + if _npu_fast_ok(a, b) and _npu_logaddexp2_fn is not None: + return _npu_logaddexp2_fn(a, b) + return _dispatch_fn("logaddexp2", None, a, b) + + r = _handle_torch_function(_py_logaddexp2_fn, (a, b), {}) + if r is not NotImplemented: + return r + + return _dispatch_fn("logaddexp2", None, a, b) + + def matmul(a, b): """Fast matmul: skip __torch_function__ when both args are base Tensor.""" cdef object r diff --git a/src/candle/_cython/_npu_ops.pyx b/src/candle/_cython/_npu_ops.pyx index d1bcbb20..205bf634 100644 --- a/src/candle/_cython/_npu_ops.pyx +++ b/src/candle/_cython/_npu_ops.pyx @@ -307,6 +307,12 @@ cdef inline void _ensure_ffi_binary() except *: global _div_getws_ptr, _div_exec_ptr global _defer_executor_fn, _acl_rt_malloc_fn, _acl_rt_free_fn global _pta_cache_begin_fn, _pta_cache_end_fn + global _atan2_getws_ptr, _atan2_exec_ptr + global _pow_tensor_tensor_getws_ptr, _pow_tensor_tensor_exec_ptr + global _remainder_getws_ptr, _remainder_exec_ptr + global _fmod_getws_ptr, _fmod_exec_ptr + global _logaddexp_getws_ptr, _logaddexp_exec_ptr + global _logaddexp2_getws_ptr, _logaddexp2_exec_ptr if _ffi_ref is not None: return from candle._cython import _aclnn_ffi as _f # pylint: disable=import-error,no-name-in-module @@ -316,6 +322,12 @@ cdef inline void _ensure_ffi_binary() except *: _mul_getws_ptr, _mul_exec_ptr = _f.resolve_op("Mul") _sub_getws_ptr, _sub_exec_ptr = _f.resolve_op("Sub") _div_getws_ptr, _div_exec_ptr = _f.resolve_op("Div") + _atan2_getws_ptr, _atan2_exec_ptr = _f.resolve_op("Atan2") + _pow_tensor_tensor_getws_ptr, _pow_tensor_tensor_exec_ptr = _f.resolve_op("PowTensorTensor") + _remainder_getws_ptr, _remainder_exec_ptr = _f.resolve_op("RemainderTensorTensor") + _fmod_getws_ptr, _fmod_exec_ptr = _f.resolve_op("FmodTensor") + _logaddexp_getws_ptr, _logaddexp_exec_ptr = _f.resolve_op("LogAddExp") + _logaddexp2_getws_ptr, _logaddexp2_exec_ptr = _f.resolve_op("LogAddExp2") _defer_executor_fn = _def_ex _acl = _eacl() _acl_rt_malloc_fn = _acl.rt.malloc @@ -836,6 +848,260 @@ def fast_div(a, b): return _cy_make_npu_tensor(out_ptr, n, a_dtype, a_dev, out_shape, out_stride) +# --------------------------------------------------------------------------- +# Shared same-dtype no-alpha binary execution helpers +# --------------------------------------------------------------------------- + +cdef _fast_binary_no_alpha_exec(a, b, object getws_ptr, object exec_ptr, str op_name): + """Shared execution helper for same-dtype no-alpha binary ops using binary_op_no_alpha FFI. + + Used by fast_atan2 and similar ops that map to aclnn binary_op_no_alpha. + Callers pass pre-resolved getws_ptr and exec_ptr (fetched via resolve_op). + No PTA cache path — straight GetWorkspaceSize + Execute. + """ + _ensure_npu_imports() + _ensure_ffi_binary() + + # 1. Validate device/dtype + a_dev = a.device + b_dev = b.device + if a_dev.type != "npu" or b_dev.type != "npu": + raise ValueError(f"fast {op_name} expects NPU tensors") + a_dtype = a.dtype + if a_dtype != b.dtype: + raise ValueError(f"fast {op_name} requires matching dtypes") + + # 2. Get runtime + stream + cdef int dev_idx = a_dev.index or 0 + runtime = _get_runtime_fast(dev_idx) + stream = _get_stream_fast(dev_idx) + + # 3. Extract shapes into C arrays + py_a_shape = a.shape + py_b_shape = b.shape + cdef int a_ndim = len(py_a_shape) + cdef int b_ndim = len(py_b_shape) + + if a_ndim > MAX_NDIM or b_ndim > MAX_NDIM: + raise ValueError(f"ndim exceeds MAX_NDIM ({MAX_NDIM})") + + cdef int64_t[MAX_NDIM] a_shape_buf, b_shape_buf + cdef int64_t[MAX_NDIM] out_shape_buf, out_stride_buf + + _fill_shape(py_a_shape, a_shape_buf, a_ndim) + _fill_shape(py_b_shape, b_shape_buf, b_ndim) + + # 4. C-level shape computation + cdef int out_ndim + cdef int64_t n + with nogil: + out_ndim = c_broadcast_shape( + a_shape_buf, a_ndim, b_shape_buf, b_ndim, out_shape_buf) + c_contiguous_stride(out_shape_buf, out_ndim, out_stride_buf) + n = c_numel(out_shape_buf, out_ndim) + + # 5. Convert to Python tuples + out_shape = _to_tuple(out_shape_buf, out_ndim) + out_stride = _to_tuple(out_stride_buf, out_ndim) + + # 6. Allocate output via cached allocator + cdef int isize = c_dtype_itemsize(a_dtype) + cdef int64_t alloc_size = n * isize + cdef object out_ptr + if dev_idx == 0: + _ensure_allocator_dev0() + out_ptr = _fast_allocator_dev0.malloc(alloc_size, stream=stream.stream) + else: + out_ptr = _get_allocator_fn_ref(dev_idx).malloc(alloc_size, stream=stream.stream) + + # 7. Get dtype code + cdef int dtype_code = _dtype_to_acl_code(a_dtype) + + # 8. Get data pointers — direct C attribute access + cdef uintptr_t a_ptr, b_ptr, o_ptr + a_ptr = a._storage._untyped._device_ptr + b_ptr = b._storage._untyped._device_ptr + o_ptr = out_ptr + + cdef uintptr_t stream_raw = int(stream.stream) + + # 9. GetWorkspaceSize + Execute + ws_size, executor = _ffi_ref.binary_op_no_alpha( + getws_ptr, exec_ptr, + py_a_shape, a.stride, + py_b_shape, b.stride, + out_shape, out_stride, + dtype_code, 2, # ACL_FORMAT_ND = 2 + a_ptr, b_ptr, o_ptr, + stream_raw) + + if ws_size: + workspace_ptr, ret = _acl_rt_malloc_fn(ws_size, 0) + if ret != 0: + raise RuntimeError(f"acl.rt.malloc failed: {ret}") + try: + ret = _ffi_ref.execute( + exec_ptr, int(workspace_ptr), ws_size, + executor, stream_raw) + if ret != 0: + raise RuntimeError(f"aclnn{op_name} execute failed: {ret}") + finally: + runtime.defer_raw_free(workspace_ptr) + + _defer_executor_fn(executor) + + return _cy_make_npu_tensor(out_ptr, n, a_dtype, a_dev, out_shape, out_stride) + + +cdef _fast_binary_two_inputs_exec(a, b, object getws_ptr, object exec_ptr, str op_name): + """Shared execution helper for same-dtype no-alpha ops using binary_two_inputs_op FFI. + + Used by fast_pow_tensor_tensor, fast_remainder, fast_fmod, fast_logaddexp, + fast_logaddexp2, and similar ops. All three dtype codes (self, other, out) + are set to the common input dtype — same-dtype only. + Callers pass pre-resolved getws_ptr and exec_ptr (fetched via resolve_op). + """ + _ensure_npu_imports() + _ensure_ffi_binary() + + # 1. Validate device/dtype + a_dev = a.device + b_dev = b.device + if a_dev.type != "npu" or b_dev.type != "npu": + raise ValueError(f"fast {op_name} expects NPU tensors") + a_dtype = a.dtype + if a_dtype != b.dtype: + raise ValueError(f"fast {op_name} requires matching dtypes") + + # 2. Get runtime + stream + cdef int dev_idx = a_dev.index or 0 + runtime = _get_runtime_fast(dev_idx) + stream = _get_stream_fast(dev_idx) + + # 3. Extract shapes into C arrays + py_a_shape = a.shape + py_b_shape = b.shape + cdef int a_ndim = len(py_a_shape) + cdef int b_ndim = len(py_b_shape) + + if a_ndim > MAX_NDIM or b_ndim > MAX_NDIM: + raise ValueError(f"ndim exceeds MAX_NDIM ({MAX_NDIM})") + + cdef int64_t[MAX_NDIM] a_shape_buf, b_shape_buf + cdef int64_t[MAX_NDIM] out_shape_buf, out_stride_buf + + _fill_shape(py_a_shape, a_shape_buf, a_ndim) + _fill_shape(py_b_shape, b_shape_buf, b_ndim) + + # 4. C-level shape computation + cdef int out_ndim + cdef int64_t n + with nogil: + out_ndim = c_broadcast_shape( + a_shape_buf, a_ndim, b_shape_buf, b_ndim, out_shape_buf) + c_contiguous_stride(out_shape_buf, out_ndim, out_stride_buf) + n = c_numel(out_shape_buf, out_ndim) + + # 5. Convert to Python tuples + out_shape = _to_tuple(out_shape_buf, out_ndim) + out_stride = _to_tuple(out_stride_buf, out_ndim) + + # 6. Allocate output via cached allocator + cdef int isize = c_dtype_itemsize(a_dtype) + cdef int64_t alloc_size = n * isize + cdef object out_ptr + if dev_idx == 0: + _ensure_allocator_dev0() + out_ptr = _fast_allocator_dev0.malloc(alloc_size, stream=stream.stream) + else: + out_ptr = _get_allocator_fn_ref(dev_idx).malloc(alloc_size, stream=stream.stream) + + # 7. Get dtype code (same for self, other, out — same-dtype only) + cdef int dtype_code = _dtype_to_acl_code(a_dtype) + + # 8. Get data pointers — direct C attribute access + cdef uintptr_t a_ptr, b_ptr, o_ptr + a_ptr = a._storage._untyped._device_ptr + b_ptr = b._storage._untyped._device_ptr + o_ptr = out_ptr + + cdef uintptr_t stream_raw = int(stream.stream) + + # 9. GetWorkspaceSize + Execute + # Pass dtype_code three times: self_dtype, other_dtype, out_dtype (same-dtype constraint) + ws_size, executor = _ffi_ref.binary_two_inputs_op( + getws_ptr, exec_ptr, + py_a_shape, a.stride, + py_b_shape, b.stride, + out_shape, out_stride, + dtype_code, dtype_code, dtype_code, 2, # ACL_FORMAT_ND = 2 + a_ptr, b_ptr, o_ptr, + stream_raw) + + if ws_size: + workspace_ptr, ret = _acl_rt_malloc_fn(ws_size, 0) + if ret != 0: + raise RuntimeError(f"acl.rt.malloc failed: {ret}") + try: + ret = _ffi_ref.execute( + exec_ptr, int(workspace_ptr), ws_size, + executor, stream_raw) + if ret != 0: + raise RuntimeError(f"aclnn{op_name} execute failed: {ret}") + finally: + runtime.defer_raw_free(workspace_ptr) + + _defer_executor_fn(executor) + + return _cy_make_npu_tensor(out_ptr, n, a_dtype, a_dev, out_shape, out_stride) + + +# BEGIN GENERATED SAME-DTYPE NO-ALPHA FAST BINARY OPS +# op=atan2 resolve=Atan2 public=atan2 +cdef object _atan2_getws_ptr = None +cdef object _atan2_exec_ptr = None + +def fast_atan2(a, b): + return _fast_binary_no_alpha_exec(a, b, _atan2_getws_ptr, _atan2_exec_ptr, "atan2") + +# op=pow_tensor_tensor resolve=PowTensorTensor public=pow +cdef object _pow_tensor_tensor_getws_ptr = None +cdef object _pow_tensor_tensor_exec_ptr = None + +def fast_pow_tensor_tensor(a, b): + return _fast_binary_two_inputs_exec(a, b, _pow_tensor_tensor_getws_ptr, _pow_tensor_tensor_exec_ptr, "pow") + +# op=remainder resolve=RemainderTensorTensor public=remainder +cdef object _remainder_getws_ptr = None +cdef object _remainder_exec_ptr = None + +def fast_remainder(a, b): + return _fast_binary_two_inputs_exec(a, b, _remainder_getws_ptr, _remainder_exec_ptr, "remainder") + +# op=fmod resolve=FmodTensor public=fmod +cdef object _fmod_getws_ptr = None +cdef object _fmod_exec_ptr = None + +def fast_fmod(a, b): + return _fast_binary_two_inputs_exec(a, b, _fmod_getws_ptr, _fmod_exec_ptr, "fmod") + +# op=logaddexp resolve=LogAddExp public=logaddexp +cdef object _logaddexp_getws_ptr = None +cdef object _logaddexp_exec_ptr = None + +def fast_logaddexp(a, b): + return _fast_binary_two_inputs_exec(a, b, _logaddexp_getws_ptr, _logaddexp_exec_ptr, "logaddexp") + +# op=logaddexp2 resolve=LogAddExp2 public=logaddexp2 +cdef object _logaddexp2_getws_ptr = None +cdef object _logaddexp2_exec_ptr = None + +def fast_logaddexp2(a, b): + return _fast_binary_two_inputs_exec(a, b, _logaddexp2_getws_ptr, _logaddexp2_exec_ptr, "logaddexp2") + +# END GENERATED SAME-DTYPE NO-ALPHA FAST BINARY OPS + + # --------------------------------------------------------------------------- # cy_npu_synchronize — fast synchronize bypassing Python dispatch overhead # --------------------------------------------------------------------------- diff --git a/src/candle/_functional.py b/src/candle/_functional.py index d1de268d..20a651fe 100644 --- a/src/candle/_functional.py +++ b/src/candle/_functional.py @@ -152,6 +152,12 @@ def _py_neg(a): mul as _cy_mul, sub as _cy_sub, div as _cy_div, + atan2 as _cy_atan2, + pow as _cy_pow, + remainder as _cy_remainder, + fmod as _cy_fmod, + logaddexp as _cy_logaddexp, + logaddexp2 as _cy_logaddexp2, matmul as _cy_matmul, relu as _cy_relu, transpose as _cy_transpose, @@ -173,6 +179,12 @@ def _py_neg(a): mul = _cy_mul sub = _cy_sub div = _cy_div + atan2 = _cy_atan2 + pow = _cy_pow + remainder = _cy_remainder + fmod = _cy_fmod + logaddexp = _cy_logaddexp + logaddexp2 = _cy_logaddexp2 matmul = _cy_matmul relu = _cy_relu transpose = _cy_transpose @@ -284,6 +296,15 @@ def frac(a): def pow(a, b): return dispatch("pow", a.device.type, a, b) +_py_pow = pow +_py_pow.__name__ = "pow" + +# Restore Cython fast-path if loaded +try: + pow = _cy_pow # noqa: F811 +except NameError: + pass + def log2(a): return dispatch("log2", a.device.type, a) @@ -437,6 +458,15 @@ def atan(a): def atan2(a, b): return dispatch("atan2", a.device.type, a, b) +_py_atan2 = atan2 +_py_atan2.__name__ = "atan2" + +# Restore Cython fast-path if it was loaded earlier +try: + atan2 = _cy_atan2 # noqa: F811 +except NameError: + pass + def asin(a): return dispatch("asin", a.device.type, a) @@ -461,10 +491,28 @@ def addcdiv(a, b, c, value=1.0): def logaddexp(a, b): return dispatch("logaddexp", a.device.type, a, b) +_py_logaddexp = logaddexp +_py_logaddexp.__name__ = "logaddexp" + +# Restore Cython fast-path if loaded +try: + logaddexp = _cy_logaddexp # noqa: F811 +except NameError: + pass + def logaddexp2(a, b): return dispatch("logaddexp2", a.device.type, a, b) +_py_logaddexp2 = logaddexp2 +_py_logaddexp2.__name__ = "logaddexp2" + +# Restore Cython fast-path if loaded +try: + logaddexp2 = _cy_logaddexp2 # noqa: F811 +except NameError: + pass + def hypot(a, b): return dispatch("hypot", a.device.type, a, b) @@ -473,10 +521,28 @@ def hypot(a, b): def remainder(a, b): return dispatch("remainder", a.device.type, a, b) +_py_remainder = remainder +_py_remainder.__name__ = "remainder" + +# Restore Cython fast-path if it was loaded earlier +try: + remainder = _cy_remainder # noqa: F811 +except NameError: + pass + def fmod(a, b): return dispatch("fmod", a.device.type, a, b) +_py_fmod = fmod +_py_fmod.__name__ = "fmod" + +# Restore Cython fast-path if it was loaded earlier +try: + fmod = _cy_fmod # noqa: F811 +except NameError: + pass + def div(a, b, *, rounding_mode=None): r = _handle_torch_function(div, (a, b), {'rounding_mode': rounding_mode}) diff --git a/tests/npu/cython/test_fast_add_hot_path.py b/tests/npu/cython/test_fast_add_hot_path.py index df405626..2a94132c 100644 --- a/tests/npu/cython/test_fast_add_hot_path.py +++ b/tests/npu/cython/test_fast_add_hot_path.py @@ -192,6 +192,344 @@ def wrapped_div(*args, **kwargs): assert calls["count"] == 0, ( f"fast_div called aclnn.div {calls['count']} time(s); expected 0" ) - assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4), ( - "fast_div output differs from expected" - ) + + +def test_fast_pow_tensor_tensor_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.rand(4, 4, device=npu_device) + 1.0 + b = torch.rand(4, 4, device=npu_device) + 0.5 + torch.npu.synchronize() + + expected = torch.pow(a, b) + torch.npu.synchronize() + + _ = torch.pow(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.pow_tensor_tensor + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "pow_tensor_tensor", wrapped) + + out = torch.pow(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_remainder_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.rand(4, 4, device=npu_device) + 1.0 + torch.npu.synchronize() + + expected = torch.remainder(a, b) + torch.npu.synchronize() + + _ = torch.remainder(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.sremainder + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "sremainder", wrapped) + + out = torch.remainder(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_fmod_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.rand(4, 4, device=npu_device) + 1.0 + torch.npu.synchronize() + + expected = torch.fmod(a, b) + torch.npu.synchronize() + + _ = torch.fmod(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.sfmod + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "sfmod", wrapped) + + out = torch.fmod(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_logaddexp_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.randn(4, 4, device=npu_device) + torch.npu.synchronize() + + expected = torch.logaddexp(a, b) + torch.npu.synchronize() + + _ = torch.logaddexp(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.slogaddexp + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "slogaddexp", wrapped) + + out = torch.logaddexp(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_logaddexp2_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.randn(4, 4, device=npu_device) + torch.npu.synchronize() + + expected = torch.logaddexp2(a, b) + torch.npu.synchronize() + + _ = torch.logaddexp2(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.slogaddexp2 + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "slogaddexp2", wrapped) + + out = torch.logaddexp2(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_pow_tensor_tensor_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.rand(4, 4, device=npu_device) + 1.0 + b = torch.rand(4, 4, device=npu_device) + 0.5 + torch.npu.synchronize() + + expected = torch.pow(a, b) + torch.npu.synchronize() + + _ = torch.pow(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.pow_tensor_tensor + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "pow_tensor_tensor", wrapped) + + out = torch.pow(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_remainder_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.rand(4, 4, device=npu_device) + 1.0 + torch.npu.synchronize() + + expected = torch.remainder(a, b) + torch.npu.synchronize() + + _ = torch.remainder(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.sremainder + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "sremainder", wrapped) + + out = torch.remainder(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_fmod_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.rand(4, 4, device=npu_device) + 1.0 + torch.npu.synchronize() + + expected = torch.fmod(a, b) + torch.npu.synchronize() + + _ = torch.fmod(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.sfmod + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "sfmod", wrapped) + + out = torch.fmod(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_logaddexp_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.randn(4, 4, device=npu_device) + torch.npu.synchronize() + + expected = torch.logaddexp(a, b) + torch.npu.synchronize() + + _ = torch.logaddexp(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.slogaddexp + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "slogaddexp", wrapped) + + out = torch.logaddexp(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_logaddexp2_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.randn(4, 4, device=npu_device) + torch.npu.synchronize() + + expected = torch.logaddexp2(a, b) + torch.npu.synchronize() + + _ = torch.logaddexp2(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.slogaddexp2 + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "slogaddexp2", wrapped) + + out = torch.logaddexp2(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) + + +def test_fast_atan2_skips_python_aclnn_wrapper(npu_device, monkeypatch): + import candle as torch + import candle._backends.npu.aclnn as aclnn_mod + import numpy as np + + a = torch.randn(4, 4, device=npu_device) + b = torch.rand(4, 4, device=npu_device) + 1.0 + torch.npu.synchronize() + + expected = torch.atan2(a, b) + torch.npu.synchronize() + + _ = torch.atan2(a, b) + torch.npu.synchronize() + + calls = {"count": 0} + original = aclnn_mod.atan2 + + def wrapped(*args, **kwargs): + calls["count"] += 1 + return original(*args, **kwargs) + + monkeypatch.setattr(aclnn_mod, "atan2", wrapped) + + out = torch.atan2(a, b) + torch.npu.synchronize() + + assert calls["count"] == 0 + assert np.allclose(out.cpu().numpy(), expected.cpu().numpy(), rtol=1e-4, atol=1e-4) diff --git a/tools/gen_npu_fast_binary_no_alpha.py b/tools/gen_npu_fast_binary_no_alpha.py new file mode 100644 index 00000000..bb75e751 --- /dev/null +++ b/tools/gen_npu_fast_binary_no_alpha.py @@ -0,0 +1,30 @@ +#!/usr/bin/env python3 +"""Generate thin same-dtype no-alpha NPU fast binary wrappers.""" + +# (func_name, resolve_op_name, public_name, helper_name) +OPS = [ + ("atan2", "Atan2", "atan2", "_fast_binary_no_alpha_exec"), + ("pow_tensor_tensor", "PowTensorTensor", "pow", "_fast_binary_two_inputs_exec"), + ("remainder", "RemainderTensorTensor", "remainder", "_fast_binary_two_inputs_exec"), + ("fmod", "FmodTensor", "fmod", "_fast_binary_two_inputs_exec"), + ("logaddexp", "LogAddExp", "logaddexp", "_fast_binary_two_inputs_exec"), + ("logaddexp2", "LogAddExp2", "logaddexp2", "_fast_binary_two_inputs_exec"), +] + + +def main() -> None: + """Print the generated wrapper block to stdout.""" + print("# BEGIN GENERATED SAME-DTYPE NO-ALPHA FAST BINARY OPS") + for op_name, resolve_name, public_name, helper_name in OPS: + print(f"# op={op_name} resolve={resolve_name} public={public_name}") + print(f"cdef object _{op_name}_getws_ptr = None") + print(f"cdef object _{op_name}_exec_ptr = None") + print() + print(f"def fast_{op_name}(a, b):") + print(f" return {helper_name}(a, b, _{op_name}_getws_ptr, _{op_name}_exec_ptr, \"{public_name}\")") + print() + print("# END GENERATED SAME-DTYPE NO-ALPHA FAST BINARY OPS") + + +if __name__ == "__main__": + main()