Skip to content

Commit 8d101ce

Browse files
committed
Unify promoter implementations and align behavior with NumPy
1 parent d927791 commit 8d101ce

7 files changed

Lines changed: 272 additions & 198 deletions

File tree

‎src/csrc/umath/binary_ops.cpp‎

Lines changed: 2 additions & 76 deletions
Original file line numberDiff line numberDiff line change
@@ -492,48 +492,10 @@ create_quad_binary_2out_ufunc(PyObject *numpy, const char *ufunc_name)
492492
return -1;
493493
}
494494

495-
PyObject *promoter_capsule =
496-
PyCapsule_New((void *)&quad_ufunc_promoter, "numpy._ufunc_promoter", NULL);
497-
if (promoter_capsule == NULL) {
495+
if (quad_add_promoters(ufunc) < 0) {
498496
Py_DECREF(ufunc);
499497
return -1;
500498
}
501-
502-
// Register promoter for (QuadPrecDType, Any, Any, Any)
503-
PyObject *DTypes = PyTuple_Pack(4, &QuadPrecDType, &PyArrayDescr_Type,
504-
&PyArrayDescr_Type, &PyArrayDescr_Type);
505-
if (DTypes == NULL) {
506-
Py_DECREF(promoter_capsule);
507-
Py_DECREF(ufunc);
508-
return -1;
509-
}
510-
511-
if (PyUFunc_AddPromoter(ufunc, DTypes, promoter_capsule) < 0) {
512-
Py_DECREF(promoter_capsule);
513-
Py_DECREF(DTypes);
514-
Py_DECREF(ufunc);
515-
return -1;
516-
}
517-
Py_DECREF(DTypes);
518-
519-
// Register promoter for (Any, QuadPrecDType, Any, Any)
520-
DTypes = PyTuple_Pack(4, &PyArrayDescr_Type, &QuadPrecDType,
521-
&PyArrayDescr_Type, &PyArrayDescr_Type);
522-
if (DTypes == NULL) {
523-
Py_DECREF(promoter_capsule);
524-
Py_DECREF(ufunc);
525-
return -1;
526-
}
527-
528-
if (PyUFunc_AddPromoter(ufunc, DTypes, promoter_capsule) < 0) {
529-
Py_DECREF(promoter_capsule);
530-
Py_DECREF(DTypes);
531-
Py_DECREF(ufunc);
532-
return -1;
533-
}
534-
Py_DECREF(promoter_capsule);
535-
Py_DECREF(DTypes);
536-
537499
Py_DECREF(ufunc);
538500
return 0;
539501
}
@@ -572,46 +534,10 @@ create_quad_binary_ufunc(PyObject *numpy, const char *ufunc_name)
572534
return -1;
573535
}
574536

575-
PyObject *promoter_capsule =
576-
PyCapsule_New((void *)&quad_ufunc_promoter, "numpy._ufunc_promoter", NULL);
577-
if (promoter_capsule == NULL) {
578-
Py_DECREF(ufunc);
579-
return -1;
580-
}
581-
582-
// Register promoter for (QuadPrecDType, Any, Any)
583-
PyObject *DTypes = PyTuple_Pack(3, &QuadPrecDType, &PyArrayDescr_Type, &PyArrayDescr_Type);
584-
if (DTypes == NULL) {
585-
Py_DECREF(promoter_capsule);
586-
Py_DECREF(ufunc);
587-
return -1;
588-
}
589-
590-
if (PyUFunc_AddPromoter(ufunc, DTypes, promoter_capsule) < 0) {
591-
Py_DECREF(promoter_capsule);
592-
Py_DECREF(DTypes);
593-
Py_DECREF(ufunc);
594-
return -1;
595-
}
596-
Py_DECREF(DTypes);
597-
598-
// Register promoter for (Any, QuadPrecDType, Any)
599-
DTypes = PyTuple_Pack(3, &PyArrayDescr_Type, &QuadPrecDType, &PyArrayDescr_Type);
600-
if (DTypes == NULL) {
601-
Py_DECREF(promoter_capsule);
602-
Py_DECREF(ufunc);
603-
return -1;
604-
}
605-
606-
if (PyUFunc_AddPromoter(ufunc, DTypes, promoter_capsule) < 0) {
607-
Py_DECREF(promoter_capsule);
608-
Py_DECREF(DTypes);
537+
if (quad_add_promoters(ufunc) < 0) {
609538
Py_DECREF(ufunc);
610539
return -1;
611540
}
612-
Py_DECREF(promoter_capsule);
613-
Py_DECREF(DTypes);
614-
615541
Py_DECREF(ufunc);
616542
return 0;
617543
}

‎src/csrc/umath/comparison_ops.cpp‎

Lines changed: 1 addition & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -269,21 +269,6 @@ comparison_object_output_is_bool(PyUFuncObject *ufunc)
269269
strcmp(ufunc->name, "logical_xor") != 0;
270270
}
271271

272-
// Registers `promoter` for a single (in1, in2, out) DType pattern.
273-
static int
274-
add_comparison_promoter(PyObject *ufunc, PyObject *promoter, PyArray_DTypeMeta *in1,
275-
PyArray_DTypeMeta *in2, PyArray_DTypeMeta *out)
276-
{
277-
PyObject *DTypes = PyTuple_Pack(3, (PyObject *)in1, (PyObject *)in2, (PyObject *)out);
278-
if (DTypes == NULL) {
279-
return -1;
280-
}
281-
282-
int res = PyUFunc_AddPromoter(ufunc, DTypes, promoter);
283-
Py_DECREF(DTypes);
284-
return res;
285-
}
286-
287272
NPY_NO_EXPORT int
288273
comparison_ufunc_promoter(PyObject *ufunc_obj, PyArray_DTypeMeta *const op_dtypes[],
289274
PyArray_DTypeMeta *const signature[], PyArray_DTypeMeta *new_op_dtypes[])
@@ -401,8 +386,7 @@ create_quad_comparison_ufunc(PyObject *numpy, const char *ufunc_name)
401386
};
402387

403388
for (size_t i = 0; i < sizeof(promoter_patterns) / sizeof(promoter_patterns[0]); i++) {
404-
if (add_comparison_promoter(ufunc, promoter_capsule, promoter_patterns[i][0],
405-
promoter_patterns[i][1], promoter_patterns[i][2]) < 0) {
389+
if (quad_add_promoter(ufunc, promoter_capsule, promoter_patterns[i]) < 0) {
406390
Py_DECREF(promoter_capsule);
407391
Py_DECREF(ufunc);
408392
return -1;

‎src/csrc/umath/matmul.cpp‎

Lines changed: 1 addition & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -518,48 +518,10 @@ init_matmul_ops(PyObject *numpy)
518518
return -1;
519519
}
520520

521-
PyObject *promoter_capsule =
522-
PyCapsule_New((void *)&quad_ufunc_promoter, "numpy._ufunc_promoter", NULL);
523-
if (promoter_capsule == NULL) {
521+
if (quad_add_promoters(ufunc) < 0) {
524522
Py_DECREF(ufunc);
525523
return -1;
526524
}
527-
528-
// Register promoter for (QuadPrecDType, Any, Any)
529-
PyObject *DTypes = PyTuple_Pack(3, &QuadPrecDType, &PyArrayDescr_Type, &PyArrayDescr_Type);
530-
if (DTypes == NULL) {
531-
Py_DECREF(promoter_capsule);
532-
Py_DECREF(ufunc);
533-
return -1;
534-
}
535-
536-
if (PyUFunc_AddPromoter(ufunc, DTypes, promoter_capsule) < 0) {
537-
Py_DECREF(promoter_capsule);
538-
Py_DECREF(DTypes);
539-
Py_DECREF(ufunc);
540-
return -1;
541-
}
542-
Py_DECREF(DTypes);
543-
544-
// Register promoter for (Any, QuadPrecDType, Any)
545-
DTypes = PyTuple_Pack(3, &PyArrayDescr_Type, &QuadPrecDType, &PyArrayDescr_Type);
546-
if (DTypes == NULL) {
547-
Py_DECREF(promoter_capsule);
548-
Py_DECREF(ufunc);
549-
return -1;
550-
}
551-
552-
if (PyUFunc_AddPromoter(ufunc, DTypes, promoter_capsule) < 0) {
553-
Py_DECREF(promoter_capsule);
554-
Py_DECREF(DTypes);
555-
Py_DECREF(ufunc);
556-
return -1;
557-
}
558-
Py_DECREF(DTypes);
559-
560-
Py_DECREF(promoter_capsule);
561-
562525
Py_DECREF(ufunc);
563-
564526
return 0;
565527
}

‎src/csrc/umath/unary_ops.cpp‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,7 @@ extern "C" {
1919
#include "scalar.h"
2020
#include "dtype.h"
2121
#include "ops.hpp"
22+
#include "umath/promoters.hpp"
2223

2324
static NPY_CASTING
2425
quad_unary_op_resolve_descriptors(PyObject *self, PyArray_DTypeMeta *const dtypes[],
@@ -140,6 +141,11 @@ create_quad_unary_ufunc(PyObject *numpy, const char *ufunc_name)
140141
return -1;
141142
}
142143

144+
if (quad_add_promoters(ufunc) < 0) {
145+
Py_DECREF(ufunc);
146+
return -1;
147+
}
148+
143149
Py_DECREF(ufunc);
144150
return 0;
145151
}
@@ -409,6 +415,11 @@ create_quad_unary_2out_ufunc(PyObject *numpy, const char *ufunc_name)
409415
return -1;
410416
}
411417

418+
if (quad_add_promoters(ufunc) < 0) {
419+
Py_DECREF(ufunc);
420+
return -1;
421+
}
422+
412423
Py_DECREF(ufunc);
413424
return 0;
414425
}

‎src/include/umath/promoters.hpp‎

Lines changed: 57 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,24 @@ quad_ufunc_has_object_input(PyUFuncObject *ufunc, PyArray_DTypeMeta *const op_dt
2222
return false;
2323
}
2424

25+
inline int
26+
quad_add_promoter(PyObject *ufunc, PyObject *promoter, PyArray_DTypeMeta *const dtypes[])
27+
{
28+
int nargs = ((PyUFuncObject *)ufunc)->nargs;
29+
PyObject *dtype_tuple = PyTuple_New(nargs);
30+
if (dtype_tuple == NULL) {
31+
return -1;
32+
}
33+
for (int i = 0; i < nargs; i++) {
34+
Py_INCREF(dtypes[i]);
35+
PyTuple_SET_ITEM(dtype_tuple, i, (PyObject *)dtypes[i]);
36+
}
37+
38+
int res = PyUFunc_AddPromoter(ufunc, dtype_tuple, promoter);
39+
Py_DECREF(dtype_tuple);
40+
return res;
41+
}
42+
2543
inline void
2644
quad_set_promoted_dtype(PyArray_DTypeMeta *signature_dtype,
2745
PyArray_DTypeMeta *fallback_dtype,
@@ -49,22 +67,51 @@ quad_ufunc_promoter(PyObject *ufunc_obj, PyArray_DTypeMeta *const op_dtypes[],
4967
return 0;
5068
}
5169

52-
if (quad_ufunc_has_object_input(ufunc, op_dtypes)) {
53-
for (int i = 0; i < nargs; i++) {
54-
quad_set_promoted_dtype(signature[i], &PyArray_ObjectDType, &new_op_dtypes[i]);
70+
PyArray_DTypeMeta *common_dtype = signature[ufunc->nin];
71+
for (int i = ufunc->nin + 1; i < nargs; i++) {
72+
if (signature[i] != common_dtype) {
73+
common_dtype = NULL;
74+
break;
5575
}
56-
return 0;
5776
}
58-
59-
// This promoter is only registered for patterns where at least one
60-
// input is QuadPrecDType, so we always promote all args to QuadPrecDType.
77+
if (common_dtype == NULL) {
78+
common_dtype = quad_ufunc_has_object_input(ufunc, op_dtypes)
79+
? &PyArray_ObjectDType : &QuadPrecDType;
80+
}
6181
for (int i = 0; i < nargs; i++) {
62-
quad_set_promoted_dtype(signature[i], &QuadPrecDType, &new_op_dtypes[i]);
82+
quad_set_promoted_dtype(signature[i], common_dtype, &new_op_dtypes[i]);
6383
}
64-
6584
return 0;
6685
}
6786

87+
inline int
88+
quad_add_promoters(PyObject *ufunc_obj)
89+
{
90+
PyUFuncObject *ufunc = (PyUFuncObject *)ufunc_obj;
91+
assert(ufunc->nin >= 1 && ufunc->nin <= 2 && ufunc->nargs <= 4);
92+
PyObject *capsule = PyCapsule_New((void *)&quad_ufunc_promoter,
93+
"numpy._ufunc_promoter", NULL);
94+
if (capsule == NULL) {
95+
return -1;
96+
}
97+
PyArray_DTypeMeta *any_dtype = (PyArray_DTypeMeta *)&PyArrayDescr_Type;
98+
PyArray_DTypeMeta *pattern[4] = {any_dtype, any_dtype, any_dtype, any_dtype};
99+
for (int i = 0; i < ufunc->nin; i++) {
100+
pattern[i] = &QuadPrecDType;
101+
}
102+
103+
// All-Quad inputs must precede mixed inputs to avoid ambiguous promotion
104+
// when an explicit output dtype excludes the Quad loop.
105+
int res = quad_add_promoter(ufunc_obj, capsule, pattern);
106+
for (int quad_slot = 0; res == 0 && ufunc->nin == 2 && quad_slot < 2; quad_slot++) {
107+
pattern[quad_slot] = &QuadPrecDType;
108+
pattern[1 - quad_slot] = any_dtype;
109+
res = quad_add_promoter(ufunc_obj, capsule, pattern);
110+
}
111+
Py_DECREF(capsule);
112+
return res;
113+
}
114+
68115

69116
inline int
70117
quad_ldexp_promoter(PyObject *ufunc_obj, PyArray_DTypeMeta *const op_dtypes[],
@@ -89,4 +136,4 @@ quad_ldexp_promoter(PyObject *ufunc_obj, PyArray_DTypeMeta *const op_dtypes[],
89136
return 0;
90137
}
91138

92-
#endif
139+
#endif

0 commit comments

Comments
 (0)