@@ -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+
2543inline void
2644quad_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
69116inline int
70117quad_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