From c4d0c9cf79053d850b3f9a2c99db6a0f044098be Mon Sep 17 00:00:00 2001 From: swayaminsync Date: Sun, 18 Jan 2026 13:53:54 +0530 Subject: [PATCH 1/2] adding argmax and argmin slots --- src/csrc/dtype.c | 166 ++++++++++++++++++++++++++++++++++++++++ tests/test_quaddtype.py | 29 +++++++ 2 files changed, 195 insertions(+) diff --git a/src/csrc/dtype.c b/src/csrc/dtype.c index 2033f01..dc47db5 100644 --- a/src/csrc/dtype.c +++ b/src/csrc/dtype.c @@ -455,6 +455,170 @@ quadprec_compare(void *a, void *b, void *arr) } } +/* + * Argmax function for np.argmax() + * Finds the index of the maximum element. + * NaN values are ignored unless all values are NaN. + */ +static int +quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) +{ + PyArrayObject *array = (PyArrayObject *)arr; + QuadPrecDTypeObject *descr = (QuadPrecDTypeObject *)PyArray_DESCR(array); + npy_intp elsize = descr->base.elsize; + + *max_ind = 0; + + if (descr->backend == BACKEND_SLEEF) { + // Find first non-NaN value as initial max + npy_intp start = 0; + for (start = 0; start < n; start++) { + Sleef_quad val = *(Sleef_quad *)(data + start * elsize); + if (!Sleef_iunordq1(val, val)) { + *max_ind = start; + break; + } + } + + // If all values are NaN, return 0 + if (start == n) { + *max_ind = 0; + return 0; + } + + // Find maximum + for (npy_intp i = start + 1; i < n; i++) { + Sleef_quad val = *(Sleef_quad *)(data + i * elsize); + Sleef_quad max_val = *(Sleef_quad *)(data + (*max_ind) * elsize); + + // Skip NaN values + if (Sleef_iunordq1(val, val)) { + continue; + } + + if (Sleef_icmpgtq1(val, max_val)) { + *max_ind = i; + } + } + } + else { + // Find first non-NaN value as initial max + npy_intp start = 0; + for (start = 0; start < n; start++) { + long double val = *(long double *)(data + start * elsize); + if (!isnan(val)) { + *max_ind = start; + break; + } + } + + // If all values are NaN, return 0 + if (start == n) { + *max_ind = 0; + return 0; + } + + // Find maximum + for (npy_intp i = start + 1; i < n; i++) { + long double val = *(long double *)(data + i * elsize); + long double max_val = *(long double *)(data + (*max_ind) * elsize); + + // Skip NaN values + if (isnan(val)) { + continue; + } + + if (val > max_val) { + *max_ind = i; + } + } + } + + return 0; +} + +/* + * Argmin function for np.argmin() + * Finds the index of the minimum element. + * NaN values are ignored unless all values are NaN. + */ +static int +quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) +{ + PyArrayObject *array = (PyArrayObject *)arr; + QuadPrecDTypeObject *descr = (QuadPrecDTypeObject *)PyArray_DESCR(array); + npy_intp elsize = descr->base.elsize; + + *min_ind = 0; + + if (descr->backend == BACKEND_SLEEF) { + // Find first non-NaN value as initial min + npy_intp start = 0; + for (start = 0; start < n; start++) { + Sleef_quad val = *(Sleef_quad *)(data + start * elsize); + if (!Sleef_iunordq1(val, val)) { + *min_ind = start; + break; + } + } + + // If all values are NaN, return 0 + if (start == n) { + *min_ind = 0; + return 0; + } + + // Find minimum + for (npy_intp i = start + 1; i < n; i++) { + Sleef_quad val = *(Sleef_quad *)(data + i * elsize); + Sleef_quad min_val = *(Sleef_quad *)(data + (*min_ind) * elsize); + + // Skip NaN values + if (Sleef_iunordq1(val, val)) { + continue; + } + + if (Sleef_icmpltq1(val, min_val)) { + *min_ind = i; + } + } + } + else { + // Find first non-NaN value as initial min + npy_intp start = 0; + for (start = 0; start < n; start++) { + long double val = *(long double *)(data + start * elsize); + if (!isnan(val)) { + *min_ind = start; + break; + } + } + + // If all values are NaN, return 0 + if (start == n) { + *min_ind = 0; + return 0; + } + + // Find minimum + for (npy_intp i = start + 1; i < n; i++) { + long double val = *(long double *)(data + i * elsize); + long double min_val = *(long double *)(data + (*min_ind) * elsize); + + // Skip NaN values + if (isnan(val)) { + continue; + } + + if (val < min_val) { + *min_ind = i; + } + } + } + + return 0; +} + static PyType_Slot QuadPrecDType_Slots[] = { {NPY_DT_ensure_canonical, &ensure_canonical}, {NPY_DT_common_instance, &common_instance}, @@ -465,6 +629,8 @@ static PyType_Slot QuadPrecDType_Slots[] = { {NPY_DT_default_descr, &quadprec_default_descr}, {NPY_DT_get_constant, &quadprec_get_constant}, {NPY_DT_PyArray_ArrFuncs_compare, &quadprec_compare}, + {NPY_DT_PyArray_ArrFuncs_argmax, &quadprec_argmax}, + {NPY_DT_PyArray_ArrFuncs_argmin, &quadprec_argmin}, {NPY_DT_PyArray_ArrFuncs_fill, &quadprec_fill}, {NPY_DT_PyArray_ArrFuncs_scanfunc, &quadprec_scanfunc}, {NPY_DT_PyArray_ArrFuncs_fromstr, &quadprec_fromstr}, diff --git a/tests/test_quaddtype.py b/tests/test_quaddtype.py index f69b944..dfab5e8 100644 --- a/tests/test_quaddtype.py +++ b/tests/test_quaddtype.py @@ -5745,3 +5745,32 @@ def test_sort_algorithms(self, backend, kind): expected = np.array([1, 2, 3, 5, 8, 9], dtype=QuadPrecDType(backend=backend)) np.testing.assert_array_equal(sorted_x, expected) + +@pytest.mark.parametrize("backend", ["sleef", "longdouble"]) +def test_argmax_argmin(backend): + """Test argmax and argmin operations.""" + # Basic integers + x = np.array([3, 1, 4, 1, 5, 9, 2, 6], dtype=QuadPrecDType(backend=backend)) + assert np.argmax(x) == 5 + assert np.argmin(x) == 1 + + # With infinity + x = np.array([1, float('inf'), 2, float('-inf'), 3], dtype=QuadPrecDType(backend=backend)) + assert np.argmax(x) == 1 # +inf is max + assert np.argmin(x) == 3 # -inf is min + + # With NaN (NaN should be ignored) + x = np.array([1, float('nan'), 5, 2], dtype=QuadPrecDType(backend=backend)) + assert np.argmax(x) == 2 + assert np.argmin(x) == 0 + + # All NaN returns index 0 + x = np.array([float('nan'), float('nan')], dtype=QuadPrecDType(backend=backend)) + assert np.argmax(x) == 0 + assert np.argmin(x) == 0 + + # 2D with axis + x = np.array([[1, 5, 3], [4, 2, 6]], dtype=QuadPrecDType(backend=backend)) + assert np.argmax(x) == 5 # flattened + np.testing.assert_array_equal(np.argmax(x, axis=0), [1, 0, 1]) + np.testing.assert_array_equal(np.argmax(x, axis=1), [1, 2]) \ No newline at end of file From d6a1028e8f6770236cba3bf2b75a36fea2e33988 Mon Sep 17 00:00:00 2001 From: swayaminsync Date: Sun, 18 Jan 2026 15:14:43 +0530 Subject: [PATCH 2/2] more tests + reviews --- src/csrc/dtype.c | 28 ++++++++++++++++------------ tests/test_quaddtype.py | 12 +++++++++++- 2 files changed, 27 insertions(+), 13 deletions(-) diff --git a/src/csrc/dtype.c b/src/csrc/dtype.c index dc47db5..414e7d4 100644 --- a/src/csrc/dtype.c +++ b/src/csrc/dtype.c @@ -472,9 +472,10 @@ quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) if (descr->backend == BACKEND_SLEEF) { // Find first non-NaN value as initial max npy_intp start = 0; + Sleef_quad max_val; for (start = 0; start < n; start++) { - Sleef_quad val = *(Sleef_quad *)(data + start * elsize); - if (!Sleef_iunordq1(val, val)) { + max_val = *(Sleef_quad *)(data + start * elsize); + if (!Sleef_iunordq1(max_val, max_val)) { *max_ind = start; break; } @@ -489,7 +490,6 @@ quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) // Find maximum for (npy_intp i = start + 1; i < n; i++) { Sleef_quad val = *(Sleef_quad *)(data + i * elsize); - Sleef_quad max_val = *(Sleef_quad *)(data + (*max_ind) * elsize); // Skip NaN values if (Sleef_iunordq1(val, val)) { @@ -497,6 +497,7 @@ quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) } if (Sleef_icmpgtq1(val, max_val)) { + max_val = val; *max_ind = i; } } @@ -504,9 +505,10 @@ quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) else { // Find first non-NaN value as initial max npy_intp start = 0; + long double max_val; for (start = 0; start < n; start++) { - long double val = *(long double *)(data + start * elsize); - if (!isnan(val)) { + max_val = *(long double *)(data + start * elsize); + if (!isnan(max_val)) { *max_ind = start; break; } @@ -521,7 +523,6 @@ quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) // Find maximum for (npy_intp i = start + 1; i < n; i++) { long double val = *(long double *)(data + i * elsize); - long double max_val = *(long double *)(data + (*max_ind) * elsize); // Skip NaN values if (isnan(val)) { @@ -529,6 +530,7 @@ quadprec_argmax(char *data, npy_intp n, npy_intp *max_ind, void *arr) } if (val > max_val) { + max_val = val; *max_ind = i; } } @@ -554,9 +556,10 @@ quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) if (descr->backend == BACKEND_SLEEF) { // Find first non-NaN value as initial min npy_intp start = 0; + Sleef_quad min_val; for (start = 0; start < n; start++) { - Sleef_quad val = *(Sleef_quad *)(data + start * elsize); - if (!Sleef_iunordq1(val, val)) { + min_val = *(Sleef_quad *)(data + start * elsize); + if (!Sleef_iunordq1(min_val, min_val)) { *min_ind = start; break; } @@ -571,7 +574,6 @@ quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) // Find minimum for (npy_intp i = start + 1; i < n; i++) { Sleef_quad val = *(Sleef_quad *)(data + i * elsize); - Sleef_quad min_val = *(Sleef_quad *)(data + (*min_ind) * elsize); // Skip NaN values if (Sleef_iunordq1(val, val)) { @@ -579,6 +581,7 @@ quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) } if (Sleef_icmpltq1(val, min_val)) { + min_val = val; *min_ind = i; } } @@ -586,9 +589,10 @@ quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) else { // Find first non-NaN value as initial min npy_intp start = 0; + long double min_val; for (start = 0; start < n; start++) { - long double val = *(long double *)(data + start * elsize); - if (!isnan(val)) { + min_val = *(long double *)(data + start * elsize); + if (!isnan(min_val)) { *min_ind = start; break; } @@ -603,7 +607,6 @@ quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) // Find minimum for (npy_intp i = start + 1; i < n; i++) { long double val = *(long double *)(data + i * elsize); - long double min_val = *(long double *)(data + (*min_ind) * elsize); // Skip NaN values if (isnan(val)) { @@ -611,6 +614,7 @@ quadprec_argmin(char *data, npy_intp n, npy_intp *min_ind, void *arr) } if (val < min_val) { + min_val = val; *min_ind = i; } } diff --git a/tests/test_quaddtype.py b/tests/test_quaddtype.py index dfab5e8..3d495b6 100644 --- a/tests/test_quaddtype.py +++ b/tests/test_quaddtype.py @@ -5772,5 +5772,15 @@ def test_argmax_argmin(backend): # 2D with axis x = np.array([[1, 5, 3], [4, 2, 6]], dtype=QuadPrecDType(backend=backend)) assert np.argmax(x) == 5 # flattened + assert np.argmin(x) == 0 # flattened np.testing.assert_array_equal(np.argmax(x, axis=0), [1, 0, 1]) - np.testing.assert_array_equal(np.argmax(x, axis=1), [1, 2]) \ No newline at end of file + np.testing.assert_array_equal(np.argmin(x, axis=0), [0, 1, 0]) + np.testing.assert_array_equal(np.argmax(x, axis=1), [1, 2]) + np.testing.assert_array_equal(np.argmin(x, axis=1), [0, 1]) + + # Empty array raises ValueError + x = np.array([], dtype=QuadPrecDType(backend=backend)) + with pytest.raises(ValueError): + np.argmax(x) + with pytest.raises(ValueError): + np.argmin(x) \ No newline at end of file