Files
petrbalvin af4ee19703
Release / gates (push) Successful in 4m38s
Test / test (push) Successful in 5m16s
Release / release (push) Successful in 35s
feat: initial release
Assisted-by: GLM 5.3 Flash
2026-09-03 10:00:00 +02:00

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)
}
}