feat: initial release
Assisted-by: GLM 5.3 Flash
This commit is contained in:
+222
@@ -0,0 +1,222 @@
|
||||
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
||||
// SPDX-License-Identifier: MIT
|
||||
|
||||
package io
|
||||
|
||||
import (
|
||||
"path/filepath"
|
||||
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestCSVRoundTrip(t *testing.T) {
|
||||
a, _ := core.FromFloats([]float64{1.5, -2.25, 3.125, 4.0, 5.5, -6.75}, 2, 3)
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "data.csv")
|
||||
if err := SaveCSV(path, a); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
back, err := LoadCSV(path, false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !core.Equal(a, back) {
|
||||
t.Errorf("CSV round-trip: got %s, want %s", back, a)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCSVWriterReader(t *testing.T) {
|
||||
a, _ := core.FromFloats([]float64{1, 2, 3, 4}, 2, 2)
|
||||
var sb strings.Builder
|
||||
if err := SaveCSVWriter(&sb, a); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if sb.String() != "1\n2\n3\n4\n" && sb.String() != "1,2\n3,4\n" {
|
||||
// csv.Writer separates with commas by default.
|
||||
if !strings.Contains(sb.String(), ",") {
|
||||
t.Fatalf("unexpected CSV: %q", sb.String())
|
||||
}
|
||||
}
|
||||
back, err := LoadCSVReader(strings.NewReader(sb.String()), false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !core.Equal(a, back) {
|
||||
t.Errorf("reader round-trip: got %s, want %s", back, a)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCSVSkipHeader(t *testing.T) {
|
||||
data := "a,b,c\n1,2,3\n4,5,6\n"
|
||||
back, err := LoadCSVReader(strings.NewReader(data), true)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
want, _ := core.FromFloats([]float64{1, 2, 3, 4, 5, 6}, 2, 3)
|
||||
if !core.Equal(back, want) {
|
||||
t.Errorf("skip header: got %s, want %s", back, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCSVErrors(t *testing.T) {
|
||||
// Non-2-D array.
|
||||
a, _ := core.FromFloats([]float64{1, 2, 3}, 3)
|
||||
if err := SaveCSV("/tmp/x.csv", a); err == nil {
|
||||
t.Error("SaveCSV: expected error for 1-D array")
|
||||
}
|
||||
// Ragged rows.
|
||||
if _, err := LoadCSVReader(strings.NewReader("1,2\n3\n"), false); err == nil {
|
||||
t.Error("LoadCSV: expected error for ragged rows")
|
||||
}
|
||||
// Non-numeric.
|
||||
if _, err := LoadCSVReader(strings.NewReader("1,x\n"), false); err == nil {
|
||||
t.Error("LoadCSV: expected error for non-numeric cell")
|
||||
}
|
||||
}
|
||||
|
||||
// TestCSVComplexRefused pins the dtype contract: CSV carries plain
|
||||
// float text, so a complex array is refused with an error instead of
|
||||
// panicking on its nil float payload.
|
||||
func TestCSVComplexRefused(t *testing.T) {
|
||||
a, err := core.FromComplexes([]complex128{complex(1, 1), complex(2, 2)}, 2, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("FromComplexes: %v", err)
|
||||
}
|
||||
var sb strings.Builder
|
||||
if err := SaveCSVWriter(&sb, a); err == nil {
|
||||
t.Error("SaveCSVWriter: expected an error for a complex array")
|
||||
}
|
||||
dir := t.TempDir()
|
||||
if err := SaveCSV(filepath.Join(dir, "c.csv"), a); err == nil {
|
||||
t.Error("SaveCSV: expected an error for a complex array")
|
||||
}
|
||||
}
|
||||
|
||||
// TestLoadCSVErrorPrecedence pins which defect a malformed file
|
||||
// reports: a record the csv parser refuses stops the read and is named
|
||||
// first, then the first row whose field count differs from the first
|
||||
// row's, then the first cell that is not a number, the order a
|
||||
// whole-file parse reports them in. The row and cell named are the
|
||||
// first offenders, never the last.
|
||||
func TestLoadCSVErrorPrecedence(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
in string
|
||||
header bool
|
||||
want string
|
||||
absent string
|
||||
}{
|
||||
{"a ragged row is named", "1,2\n3\n", false, "row 1 has 1 fields, want 2", ""},
|
||||
{"the first ragged row is named", "1,2\n3\n4,5,6\n", false, "row 1 has 1 fields, want 2", "row 2"},
|
||||
{"the first bad cell is named", "1,2\n3,x\n4,y\n", false, "row 1 col 1", "row 2"},
|
||||
{"a malformed record outranks a ragged row", "1,2\n3\n\"unterminated\n", false, "parse error", "fields, want"},
|
||||
{"a ragged row outranks an earlier bad cell", "1,2\nx,2\n3\n", false, "row 2 has 1 fields, want 2", "col"},
|
||||
{"the header is not a data row", "a,b\n1,2\n3\n", true, "row 1 has 1 fields, want 2", ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
_, err := LoadCSVReader(strings.NewReader(tc.in), tc.header)
|
||||
if err == nil {
|
||||
t.Fatalf("%s: %q parsed without an error", tc.name, tc.in)
|
||||
}
|
||||
if !strings.Contains(err.Error(), tc.want) {
|
||||
t.Fatalf("%s: %v, want the message to carry %q", tc.name, err, tc.want)
|
||||
}
|
||||
if tc.absent != "" && strings.Contains(err.Error(), tc.absent) {
|
||||
t.Fatalf("%s: %v, want no mention of %q", tc.name, err, tc.absent)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestCSVWriterDtypes pins the writer's text form for every dtype the
|
||||
// core carries: the integer class through exact decimal widenings (an
|
||||
// int64 past 2^53 keeps every digit, which a float64 detour would
|
||||
// round away), Bool as 0 and 1, the numeric column semantics CSV
|
||||
// carries, and Float16 through its exact float64 widening in the same
|
||||
// 'g' form the float paths use.
|
||||
func TestCSVWriterDtypes(t *testing.T) {
|
||||
text := func(t *testing.T, a *core.Array) string {
|
||||
t.Helper()
|
||||
var sb strings.Builder
|
||||
if err := SaveCSVWriter(&sb, a); err != nil {
|
||||
t.Fatalf("SaveCSVWriter(%s): %v", a.Dtype(), err)
|
||||
}
|
||||
return sb.String()
|
||||
}
|
||||
bools, err := core.FromBools([]bool{true, false, false, true}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromBools: %v", err)
|
||||
}
|
||||
if got, want := text(t, bools), "1,0\n0,1\n"; got != want {
|
||||
t.Fatalf("bool wrote %q, want %q", got, want)
|
||||
}
|
||||
i8, err := core.FromInt8s([]int8{-128, 127, 1, -1}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInt8s: %v", err)
|
||||
}
|
||||
if got, want := text(t, i8), "-128,127\n1,-1\n"; got != want {
|
||||
t.Fatalf("int8 wrote %q, want %q", got, want)
|
||||
}
|
||||
u8, err := core.FromUint8s([]uint8{0, 255, 10, 42}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromUint8s: %v", err)
|
||||
}
|
||||
if got, want := text(t, u8), "0,255\n10,42\n"; got != want {
|
||||
t.Fatalf("uint8 wrote %q, want %q", got, want)
|
||||
}
|
||||
i16, err := core.FromInt16s([]int16{-32768, 32767, 0, -7}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInt16s: %v", err)
|
||||
}
|
||||
if got, want := text(t, i16), "-32768,32767\n0,-7\n"; got != want {
|
||||
t.Fatalf("int16 wrote %q, want %q", got, want)
|
||||
}
|
||||
u16, err := core.FromUint16s([]uint16{0, 65535, 5, 4096}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromUint16s: %v", err)
|
||||
}
|
||||
if got, want := text(t, u16), "0,65535\n5,4096\n"; got != want {
|
||||
t.Fatalf("uint16 wrote %q, want %q", got, want)
|
||||
}
|
||||
i32, err := core.FromInt32s([]int32{-2147483648, 2147483647, 0, -1}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInt32s: %v", err)
|
||||
}
|
||||
if got, want := text(t, i32), "-2147483648,2147483647\n0,-1\n"; got != want {
|
||||
t.Fatalf("int32 wrote %q, want %q", got, want)
|
||||
}
|
||||
u32, err := core.FromUint32s([]uint32{0, 4294967295, 7, 65536}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromUint32s: %v", err)
|
||||
}
|
||||
if got, want := text(t, u32), "0,4294967295\n7,65536\n"; got != want {
|
||||
t.Fatalf("uint32 wrote %q, want %q", got, want)
|
||||
}
|
||||
// The %d contract in its element: 2^53+1 cannot survive a float64
|
||||
// detour, and the integer path must not take one.
|
||||
big, err := core.FromInts([]int64{9007199254740993, -9007199254740993, 0, 2}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromInts: %v", err)
|
||||
}
|
||||
if got, want := text(t, big), "9007199254740993,-9007199254740993\n0,2\n"; got != want {
|
||||
t.Fatalf("int wrote %q, want %q", got, want)
|
||||
}
|
||||
half, err := core.FromFloat16s([]float64{0.5, -2, 1, 4}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat16s: %v", err)
|
||||
}
|
||||
if got, want := text(t, half), "0.5,-2\n1,4\n"; got != want {
|
||||
t.Fatalf("float16 wrote %q, want %q", got, want)
|
||||
}
|
||||
f32, err := core.FromFloat32s([]float32{1.5, -2.25, 0, 3.25}, 2, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("FromFloat32s: %v", err)
|
||||
}
|
||||
if got, want := text(t, f32), "1.5,-2.25\n0,3.25\n"; got != want {
|
||||
t.Fatalf("float32 wrote %q, want %q", got, want)
|
||||
}
|
||||
f64 := mustFloats(t, []float64{1.5, -2.25, 0, 4}, 2, 2)
|
||||
if got, want := text(t, f64), "1.5,-2.25\n0,4\n"; got != want {
|
||||
t.Fatalf("float wrote %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user