270 lines
7.9 KiB
Go
270 lines
7.9 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|||
|
|
// SPDX-License-Identifier: MIT
|
||
|
|
|
||
|
|
package core
|
||
|
|
|
||
|
|
import (
|
||
|
|
"strings"
|
||
|
|
"testing"
|
||
|
|
)
|
||
|
|
|
||
|
|
func mustFromComplexes(t *testing.T, vals []complex128, shape ...int) *Array {
|
||
|
|
t.Helper()
|
||
|
|
a, err := FromComplexes(vals, shape...)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("FromComplexes(%v, %v): %v", vals, shape, err)
|
||
|
|
}
|
||
|
|
return a
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexConstructors(t *testing.T) {
|
||
|
|
a := mustFromComplexes(t, []complex128{complex(1, 2), complex(3, -4)}, 2)
|
||
|
|
if a.Dtype() != Complex || a.Len() != 2 {
|
||
|
|
t.Fatalf("complex array: %s len %d", a.Dtype(), a.Len())
|
||
|
|
}
|
||
|
|
if v, err := ComplexAt(a, 0); err != nil || v != complex(1, 2) {
|
||
|
|
t.Fatalf("ComplexAt: %v %v", v, err)
|
||
|
|
}
|
||
|
|
if _, err := IntAt(a, 0); err == nil || !strings.Contains(err.Error(), "not int") {
|
||
|
|
t.Fatalf("IntAt on complex: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := FloatAt(a, 0); err == nil || !strings.Contains(err.Error(), "not float") {
|
||
|
|
t.Fatalf("FloatAt on complex: %v", err)
|
||
|
|
}
|
||
|
|
|
||
|
|
z, err := Zeros(Complex, 2)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Zeros complex: %v", err)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(z, 1); v != 0 {
|
||
|
|
t.Fatalf("Zeros complex value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
one, _ := Ones(Complex, 2)
|
||
|
|
if v, _ := ComplexAt(one, 0); v != 1 {
|
||
|
|
t.Fatalf("Ones complex value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
full, _ := FullC(complex(0.5, -0.5), 2)
|
||
|
|
if v, _ := ComplexAt(full, 1); v != complex(0.5, -0.5) {
|
||
|
|
t.Fatalf("FullC value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
id, _ := Identity(Complex, 2)
|
||
|
|
if v, _ := ComplexAt(id, 1, 1); v != 1 {
|
||
|
|
t.Fatalf("Identity complex: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(id, 0, 1); v != 0 {
|
||
|
|
t.Fatalf("Identity complex off-diagonal: %v", v)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexPromotion(t *testing.T) {
|
||
|
|
c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1)
|
||
|
|
i := mustFromInts(t, []int64{2}, 1)
|
||
|
|
f := mustFromFloats(t, []float64{0.5}, 1)
|
||
|
|
|
||
|
|
sum, err := Add(c, i)
|
||
|
|
if err != nil || sum.Dtype() != Complex {
|
||
|
|
t.Fatalf("complex + int: %s %v", sum.Dtype(), err)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(sum, 0); v != complex(3, 1) {
|
||
|
|
t.Fatalf("complex + int value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
// float + complex promotes too.
|
||
|
|
sum2, _ := Add(f, c)
|
||
|
|
if sum2.Dtype() != Complex {
|
||
|
|
t.Fatalf("float + complex dtype: %s", sum2.Dtype())
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(sum2, 0); v != complex(1.5, 1) {
|
||
|
|
t.Fatalf("float + complex value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Division on complex stays complex.
|
||
|
|
q, _ := Div(c, f)
|
||
|
|
if q.Dtype() != Complex {
|
||
|
|
t.Fatalf("complex division dtype: %s", q.Dtype())
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(q, 0); v != complex(1, 1)/complex(0.5, 0) {
|
||
|
|
t.Fatalf("complex division value: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Quo rejects complex.
|
||
|
|
if _, err := Quo(c, c); err == nil || !strings.Contains(err.Error(), "needs int arrays") {
|
||
|
|
t.Fatalf("Quo complex: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexScalarMirrors(t *testing.T) {
|
||
|
|
c := mustFromComplexes(t, []complex128{complex(1, 2)}, 1)
|
||
|
|
i := mustFromInts(t, []int64{1}, 1)
|
||
|
|
|
||
|
|
// Int scalars keep complex complex.
|
||
|
|
if v, _ := ComplexAt(AddI(c, 5), 0); v != complex(6, 2) {
|
||
|
|
t.Fatalf("AddI on complex: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(MulI(c, 2), 0); v != complex(2, 4) {
|
||
|
|
t.Fatalf("MulI on complex: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(DivF(c, 2), 0); v != complex(0.5, 1) {
|
||
|
|
t.Fatalf("DivF on complex: %v", v)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Complex scalars promote everything.
|
||
|
|
if v, _ := ComplexAt(AddC(i, complex(0, 1)), 0); v != complex(1, 1) {
|
||
|
|
t.Fatalf("AddC on int: %v", v)
|
||
|
|
}
|
||
|
|
out := MulC(i, complex(0, 2))
|
||
|
|
if out.Dtype() != Complex {
|
||
|
|
t.Fatalf("MulC dtype: %s", out.Dtype())
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(out, 0); v != complex(0, 2) {
|
||
|
|
t.Fatalf("MulC value: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(SubC(c, complex(1, 2)), 0); v != 0 {
|
||
|
|
t.Fatalf("SubC value: %v", v)
|
||
|
|
}
|
||
|
|
if v, _ := ComplexAt(DivC(c, complex(0, 1)), 0); v != complex(2, -1) {
|
||
|
|
t.Fatalf("DivC value: %v", v)
|
||
|
|
}
|
||
|
|
if _, err := QuoI(i, 1); err != nil {
|
||
|
|
t.Fatalf("QuoI still works on int: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexSumAndDot(t *testing.T) {
|
||
|
|
c := mustFromComplexes(t, []complex128{complex(1, 1), complex(2, -1)}, 2)
|
||
|
|
s := Sum(c)
|
||
|
|
if !s.IsComplex() || s.Complex() != complex(3, 0) {
|
||
|
|
t.Fatalf("Sum complex: %s", s)
|
||
|
|
}
|
||
|
|
|
||
|
|
d, err := Dot(c, c)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Dot complex: %v", err)
|
||
|
|
}
|
||
|
|
// (1+i)(1+i) + (2-i)(2-i) = 2i + 3-4i = 3-2i
|
||
|
|
if !d.IsComplex() || d.Complex() != complex(3, -2) {
|
||
|
|
t.Fatalf("Dot complex: %s", d)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Mixed Dot promotes to complex.
|
||
|
|
dm, _ := Dot(mustFromInts(t, []int64{1, 1}, 2), c)
|
||
|
|
if !dm.IsComplex() {
|
||
|
|
t.Fatalf("Dot mixed: %s", dm)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexUnorderedErrors(t *testing.T) {
|
||
|
|
c := mustFromComplexes(t, []complex128{complex(1, 1)}, 1)
|
||
|
|
|
||
|
|
if _, err := Min(c); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
||
|
|
t.Fatalf("Min complex: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := Max(c); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
||
|
|
t.Fatalf("Max complex: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := Mean(c); err == nil || !strings.Contains(err.Error(), "no float mean") {
|
||
|
|
t.Fatalf("Mean complex: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := Lt(c, c); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
||
|
|
t.Fatalf("Lt complex: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := GtI(c, 1); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
||
|
|
t.Fatalf("GtI complex: %v", err)
|
||
|
|
}
|
||
|
|
if _, err := LtF(c, 1); err == nil || !strings.Contains(err.Error(), "no ordering") {
|
||
|
|
t.Fatalf("LtF complex: %v", err)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexEqNe(t *testing.T) {
|
||
|
|
a := mustFromComplexes(t, []complex128{complex(1, 2), complex(3, 4)}, 2)
|
||
|
|
b := mustFromComplexes(t, []complex128{complex(1, 2), complex(4, 4)}, 2)
|
||
|
|
|
||
|
|
eq, err := Eq(a, b)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("Eq complex: %v", err)
|
||
|
|
}
|
||
|
|
if got := maskOf(t, eq); got[0] != 1 || got[1] != 0 {
|
||
|
|
t.Fatalf("Eq complex: %v", got)
|
||
|
|
}
|
||
|
|
ne, _ := Ne(a, b)
|
||
|
|
if got := maskOf(t, ne); got[0] != 0 || got[1] != 1 {
|
||
|
|
t.Fatalf("Ne complex: %v", got)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Complex vs float compares in complex space; only the element whose
|
||
|
|
// imaginary part is zero can match.
|
||
|
|
realish := mustFromComplexes(t, []complex128{1, complex(3, 4)}, 2)
|
||
|
|
f := mustFromFloats(t, []float64{1, 4}, 2)
|
||
|
|
eqf, _ := Eq(realish, f)
|
||
|
|
if got := maskOf(t, eqf); got[0] != 1 || got[1] != 0 {
|
||
|
|
t.Fatalf("Eq complex-float: %v", got)
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
func TestComplexMatMulAndMachinery(t *testing.T) {
|
||
|
|
m := mustFromComplexes(t, []complex128{complex(0, 1), 0, 0, complex(0, 1)}, 2, 2)
|
||
|
|
v := mustFromComplexes(t, []complex128{1, 1}, 2, 1)
|
||
|
|
|
||
|
|
out, err := MatMul2D(m, v)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("MatMul complex: %v", err)
|
||
|
|
}
|
||
|
|
if out.Dtype() != Complex || out.Shape()[0] != 2 {
|
||
|
|
t.Fatalf("MatMul complex shape: %s", out)
|
||
|
|
}
|
||
|
|
if w, _ := ComplexAt(out, 0, 0); w != complex(0, 1) {
|
||
|
|
t.Fatalf("MatMul complex value: %v", w)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Transpose, Slice, Row, Mask and WithComplex all carry the dtype.
|
||
|
|
tt := Transpose(m)
|
||
|
|
if tt.Dtype() != Complex {
|
||
|
|
t.Fatalf("Transpose complex dtype: %s", tt.Dtype())
|
||
|
|
}
|
||
|
|
s, _ := Slice(m, 0, 0, 1)
|
||
|
|
if s.Dtype() != Complex || s.Shape()[0] != 1 {
|
||
|
|
t.Fatalf("Slice complex: %s", s)
|
||
|
|
}
|
||
|
|
r, _ := Row(m, 1)
|
||
|
|
if v2, _ := ComplexAt(r, 1); v2 != complex(0, 1) {
|
||
|
|
t.Fatalf("Row complex: %v", v2)
|
||
|
|
}
|
||
|
|
upd, err := WithComplex(m, complex(9, 9), 0, 0)
|
||
|
|
if err != nil {
|
||
|
|
t.Fatalf("WithComplex: %v", err)
|
||
|
|
}
|
||
|
|
if v2, _ := ComplexAt(upd, 0, 0); v2 != complex(9, 9) {
|
||
|
|
t.Fatalf("WithComplex value: %v", v2)
|
||
|
|
}
|
||
|
|
if v2, _ := ComplexAt(m, 0, 0); v2 != complex(0, 1) {
|
||
|
|
t.Fatalf("WithComplex receiver mutated: %v", v2)
|
||
|
|
}
|
||
|
|
|
||
|
|
// Where promotes to complex and Mask selects complex elements.
|
||
|
|
cond := mustFromInts(t, []int64{1, 0}, 2)
|
||
|
|
w, _ := Where(cond, mustFromComplexes(t, []complex128{1, 1}, 2), mustFromComplexes(t, []complex128{0, 0}, 2))
|
||
|
|
if w.Dtype() != Complex {
|
||
|
|
t.Fatalf("Where complex dtype: %s", w.Dtype())
|
||
|
|
}
|
||
|
|
if v2, _ := ComplexAt(w, 1); v2 != 0 {
|
||
|
|
t.Fatalf("Where complex value: %v", v2)
|
||
|
|
}
|
||
|
|
mask := cond
|
||
|
|
sel, _ := Select(mustFromComplexes(t, []complex128{complex(5, 5), complex(6, 6)}, 2), mask)
|
||
|
|
if sel.Len() != 1 {
|
||
|
|
t.Fatalf("Mask complex len: %d", sel.Len())
|
||
|
|
}
|
||
|
|
if v2, _ := ComplexAt(sel, 0); v2 != complex(5, 5) {
|
||
|
|
t.Fatalf("Mask complex value: %v", v2)
|
||
|
|
}
|
||
|
|
|
||
|
|
// String renders complex elements.
|
||
|
|
if got := mustFromComplexes(t, []complex128{complex(1, 2)}, 1).String(); !strings.Contains(got, "(1+2i)") {
|
||
|
|
t.Fatalf("String complex: %q", got)
|
||
|
|
}
|
||
|
|
}
|