From 8071831763f33956faddd7e656b3d11bcc65e926 Mon Sep 17 00:00:00 2001 From: mrava87 Date: Thu, 13 Aug 2026 08:35:48 +0000 Subject: [PATCH] fix: modify min-max in _Simplex_cuda due to bug in numba-cuda --- pyproximal/proximal/_Simplex_cuda.py | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/pyproximal/proximal/_Simplex_cuda.py b/pyproximal/proximal/_Simplex_cuda.py index 87ec39c..af36a71 100644 --- a/pyproximal/proximal/_Simplex_cuda.py +++ b/pyproximal/proximal/_Simplex_cuda.py @@ -2,12 +2,27 @@ from numba import cuda +@cuda.jit(device=True) +def clamp_jit_cuda(v, lower, upper): + """Clamp a value between lower and upper bounds + + Note: equivalent to ``min(max(v, lower), upper)``, written out explicitly + to avoid a Numba CUDA-target typing bug in the variadic ``min``/``max`` + builtin overloads (raises a spurious "Signature mismatch" TypeError). + """ + if v < lower: + return lower + elif v > upper: + return upper + return v + + @cuda.jit(device=True) def fun_jit_cuda(mu, x, coeffs, scalar, lower, upper): """Bisection function""" p = 0 for i in range(coeffs.shape[0]): - p += coeffs[i] * min(max(x[i] - mu * coeffs[i], lower), upper) + p += coeffs[i] * clamp_jit_cuda(x[i] - mu * coeffs[i], lower, upper) return p - scalar @@ -83,4 +98,4 @@ def simplex_jit_cuda(x, coeffs, scalar, lower, upper, maxiter, ftol, xtol, y): ) for j in range(coeffs.shape[0]): - y[i][j] = min(max(x[i][j] - c * coeffs[j], lower), upper) + y[i][j] = clamp_jit_cuda(x[i][j] - c * coeffs[j], lower, upper)