Skip to content

Commit fec2bea

Browse files
committed
Add a conjugate implementation for the scalar
1 parent 8d101ce commit fec2bea

2 files changed

Lines changed: 83 additions & 0 deletions

File tree

‎src/csrc/scalar.c‎

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -429,6 +429,34 @@ QuadPrecision_get_imag(QuadPrecisionObject *self, void *closure)
429429
return (PyObject *)QuadPrecision_raw_new(self->backend);
430430
}
431431

432+
static PyObject *
433+
QuadPrecision_conjugate(QuadPrecisionObject *self, PyObject *args)
434+
{
435+
PyArrayObject *out = NULL;
436+
if (!PyArg_ParseTuple(args, "|O&:conjugate", PyArray_OutputConverter, &out)) {
437+
return NULL;
438+
}
439+
if (out == NULL) {
440+
return Py_NewRef(self);
441+
}
442+
443+
PyArray_Descr *dtype = (PyArray_Descr *)new_quaddtype_instance(self->backend);
444+
if (dtype == NULL) {
445+
return NULL;
446+
}
447+
PyArrayObject *array = (PyArrayObject *)PyArray_SimpleNewFromDescr(0, NULL, dtype);
448+
if (array == NULL) {
449+
return NULL;
450+
}
451+
quad_value_store(PyArray_BYTES(array), &self->value, self->backend);
452+
PyObject *result = PyArray_Conjugate(array, out);
453+
Py_DECREF(array);
454+
if (result == NULL || !PyArray_Check(result)) {
455+
return result;
456+
}
457+
return PyArray_Return((PyArrayObject *)result);
458+
}
459+
432460
// Method implementations for float compatibility
433461
static PyObject *
434462
QuadPrecision_is_integer(QuadPrecisionObject *self, PyObject *Py_UNUSED(ignored))
@@ -774,6 +802,10 @@ QuadPrecision_from_raw_bytes(PyObject *Py_UNUSED(module), PyObject *args)
774802
}
775803

776804
static PyMethodDef QuadPrecision_methods[] = {
805+
{"conj", (PyCFunction)QuadPrecision_conjugate, METH_VARARGS,
806+
"Return the complex conjugate."},
807+
{"conjugate", (PyCFunction)QuadPrecision_conjugate, METH_VARARGS,
808+
"Return the complex conjugate."},
777809
{"is_integer", (PyCFunction)QuadPrecision_is_integer, METH_NOARGS,
778810
"Return True if the value is an integer."},
779811
{"as_integer_ratio", (PyCFunction)QuadPrecision_as_integer_ratio, METH_NOARGS,

‎tests/test_quaddtype.py‎

Lines changed: 51 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3697,6 +3697,57 @@ def test_array_conjugate_method():
36973697
np.testing.assert_array_equal(result.astype(float), [1.5, -2.5, 0.0])
36983698

36993699

3700+
@pytest.mark.parametrize("backend", ["sleef", "longdouble"])
3701+
def test_object_conjugate_preserves_values_and_backend(backend):
3702+
dtype = QuadPrecDType(backend=backend)
3703+
objects = np.array([1.5, 2.5], dtype=dtype).astype(object)
3704+
3705+
result = np.conjugate(objects)
3706+
3707+
np.testing.assert_array_equal(result, objects, strict=True)
3708+
for scalar in result:
3709+
assert scalar.dtype == dtype
3710+
3711+
3712+
@pytest.mark.parametrize("backend", ["sleef", "longdouble"])
3713+
@pytest.mark.parametrize("method", ["conj", "conjugate"])
3714+
@pytest.mark.parametrize("value", [1.5, -0.0, np.inf, np.nan])
3715+
@pytest.mark.parametrize("args", [(), (None,)])
3716+
def test_scalar_conjugate_preserves_values_and_backend(backend, method, value, args):
3717+
scalar = QuadPrecision(value, backend=backend)
3718+
3719+
result = getattr(scalar, method)(*args)
3720+
3721+
assert result.dtype == scalar.dtype
3722+
np.testing.assert_array_equal(float(result), value)
3723+
assert np.signbit(float(result)) == np.signbit(value)
3724+
3725+
3726+
@pytest.mark.parametrize("backend", ["sleef", "longdouble"])
3727+
@pytest.mark.parametrize("method", ["conj", "conjugate"])
3728+
@pytest.mark.parametrize("dtype", [np.float32, np.float64, object, QuadPrecDType])
3729+
def test_scalar_conjugate_with_out(backend, method, dtype):
3730+
scalar = QuadPrecision(1.5, backend=backend)
3731+
out = np.empty((), dtype=scalar.dtype if dtype is QuadPrecDType else dtype)
3732+
reference_out = np.empty((), dtype=np.float64 if dtype is QuadPrecDType else dtype)
3733+
3734+
result = getattr(scalar, method)(out)
3735+
expected = getattr(np.float64(1.5), method)(reference_out)
3736+
3737+
np.testing.assert_array_equal(float(result), float(expected))
3738+
np.testing.assert_array_equal(out.astype(np.float64), reference_out.astype(np.float64))
3739+
if dtype is object or dtype is QuadPrecDType:
3740+
assert result.dtype == scalar.dtype
3741+
3742+
3743+
@pytest.mark.parametrize("backend", ["sleef", "longdouble"])
3744+
@pytest.mark.parametrize("method", ["conj", "conjugate"])
3745+
@pytest.mark.parametrize("args", [(1,), (None, None), (np.empty((), dtype=np.int64),)])
3746+
def test_scalar_conjugate_rejects_invalid_out(backend, method, args):
3747+
for scalar in [QuadPrecision(1.5, backend=backend), np.float64(1.5)]:
3748+
with pytest.raises(TypeError):
3749+
getattr(scalar, method)(*args)
3750+
37003751
@pytest.mark.parametrize("x1,x2,expected", [
37013752
# Basic Pythagorean triples
37023753
(3.0, 4.0, 5.0),

0 commit comments

Comments
 (0)