472 lines
15 KiB
Go
472 lines
15 KiB
Go
// Copyright (c) 2026 Petr Balvín <opensource@petrbalvin.org> (https://petrbalvin.org)
|
|
// SPDX-License-Identifier: MIT
|
|
|
|
package io
|
|
|
|
import (
|
|
"bufio"
|
|
"bytes"
|
|
"encoding/csv"
|
|
"io"
|
|
"os"
|
|
"strconv"
|
|
"unsafe"
|
|
|
|
"sourcedock.dev/petrbalvin/tensor/internal/base"
|
|
"sourcedock.dev/petrbalvin/tensor/internal/core"
|
|
)
|
|
|
|
// CSV IO. SaveCSV writes a 2-D array as comma-separated text;
|
|
// LoadCSV reads it back. CSV is the interchange format for tabular
|
|
// data with spreadsheets and statistical packages, and the plain-text
|
|
// counterpart to the binary FITS format for data pipelines. Only 2-D
|
|
// arrays are supported (rows by columns), matching the tabular model.
|
|
// The writer formats every dtype the core carries: the integer class
|
|
// as exact integer text, booleans as 0 and 1, floats as float text;
|
|
// the reader parses everything back as float64.
|
|
|
|
// SaveCSV writes a 2-D array to path as comma-separated values. The
|
|
// result is loadable by any spreadsheet or data tool. The close error
|
|
// is part of the write: os.Create's descriptor buffers nothing itself,
|
|
// but the operating system still reports a failed write through Close
|
|
// on some filesystems, so a discarded Close error would report a save
|
|
// that never reached the disk.
|
|
func SaveCSV(path string, a *core.Array) error {
|
|
if a.NDim() != 2 {
|
|
return base.Errf("SaveCSV: needs a 2-D array, got shape %s", base.ShapeText(a.Shape()))
|
|
}
|
|
f, err := os.Create(path)
|
|
if err != nil {
|
|
return base.Errf("SaveCSV: %w", err)
|
|
}
|
|
werr := SaveCSVWriter(f, a)
|
|
if cerr := f.Close(); werr != nil {
|
|
return base.Errf("SaveCSV: %w", werr)
|
|
} else if cerr != nil {
|
|
return base.Errf("SaveCSV: %w", cerr)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// SaveCSVWriter writes a 2-D array as CSV to w. Every dtype the core
|
|
// carries has a text form: Bool as 0 and 1, the integer dtypes as
|
|
// exact decimal widenings of their payload, Float16 through its exact
|
|
// float64 widening, and Float32 and Float as 'g' float text. Complex
|
|
// arrays are refused: CSV carries plain numeric text, and a pair of
|
|
// raw halves would read back as two unrelated columns. An array with
|
|
// no columns is refused for the same reason: one empty record per row
|
|
// is what the writer would emit, and CSV readers treat blank lines as
|
|
// no records at all, so the shape would come back as (0,0).
|
|
func SaveCSVWriter(w io.Writer, a *core.Array) error {
|
|
if a.NDim() != 2 {
|
|
return base.Errf("SaveCSV: needs a 2-D array, got shape %s", base.ShapeText(a.Shape()))
|
|
}
|
|
if a.Dtype() == core.Complex {
|
|
return base.Errf("SaveCSV: complex arrays are not supported")
|
|
}
|
|
cw := csv.NewWriter(w)
|
|
rows, cols := a.Shape()[0], a.Shape()[1]
|
|
if cols == 0 {
|
|
// No column count can follow from CSV text that has no record.
|
|
return base.Errf("SaveCSV: an array with no columns has no CSV form, got shape %s", base.ShapeText(a.Shape()))
|
|
}
|
|
// The package's arrays, views included, always keep payload[i] at
|
|
// element i (views rebase the payload, never stride it), so the
|
|
// row walk indexes the payload directly instead of going through
|
|
// the per-element accessor.
|
|
rec := make([]string, cols)
|
|
for r := range rows {
|
|
baseIdx := r * cols
|
|
switch a.Dtype() {
|
|
case core.Float:
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatFloat(a.RawFloats()[baseIdx+c], 'g', -1, 64)
|
|
}
|
|
case core.Float32:
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatFloat(float64(a.RawFloat32s()[baseIdx+c]), 'g', -1, 64)
|
|
}
|
|
case core.Float16:
|
|
// The half payload widens to float64 exactly, and the
|
|
// widened value formats the way every float path formats.
|
|
vals := a.RawHalves()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatFloat(core.HalfToFloat64(vals[baseIdx+c]), 'g', -1, 64)
|
|
}
|
|
case core.Bool:
|
|
// Bool writes as 0 and 1: the numeric column semantics a
|
|
// CSV column carries, the two values the payload stores.
|
|
vals := a.RawBools()
|
|
for c := range cols {
|
|
if vals[baseIdx+c] {
|
|
rec[c] = "1"
|
|
} else {
|
|
rec[c] = "0"
|
|
}
|
|
}
|
|
case core.Int8:
|
|
vals := a.RawInt8s()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10)
|
|
}
|
|
case core.Uint8:
|
|
vals := a.RawUint8s()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10)
|
|
}
|
|
case core.Int16:
|
|
vals := a.RawInt16s()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10)
|
|
}
|
|
case core.Uint16:
|
|
vals := a.RawUint16s()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10)
|
|
}
|
|
case core.Int32:
|
|
vals := a.RawInt32s()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10)
|
|
}
|
|
case core.Uint32:
|
|
vals := a.RawUint32s()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(int64(vals[baseIdx+c]), 10)
|
|
}
|
|
case core.Int:
|
|
// The %d contract: an int64 past 2^53 keeps every digit,
|
|
// where a float64 detour would round it away.
|
|
vals := a.RawInts()
|
|
for c := range cols {
|
|
rec[c] = strconv.FormatInt(vals[baseIdx+c], 10)
|
|
}
|
|
default:
|
|
// Unreachable: Complex is refused above and every other
|
|
// dtype has its case; the writer answers an error rather
|
|
// than panic on a payload it cannot name.
|
|
return base.Errf("SaveCSV: dtype %s has no CSV form", a.Dtype())
|
|
}
|
|
if err := cw.Write(rec); err != nil {
|
|
return base.Errf("SaveCSV: %w", err)
|
|
}
|
|
}
|
|
cw.Flush()
|
|
if err := cw.Error(); err != nil {
|
|
return base.Errf("SaveCSV: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// LoadCSV reads a 2-D float array from a CSV file. Every row must have
|
|
// the same number of fields; a header row is treated as data unless
|
|
// skipHeader is true.
|
|
func LoadCSV(path string, skipHeader bool) (*core.Array, error) {
|
|
f, err := os.Open(path)
|
|
if err != nil {
|
|
return nil, base.Errf("LoadCSV: %w", err)
|
|
}
|
|
defer f.Close()
|
|
return LoadCSVReader(f, skipHeader)
|
|
}
|
|
|
|
// LoadCSVReader reads a 2-D float array from a CSV stream. A leading
|
|
// UTF-8 byte order mark is skipped: spreadsheets write one, and the
|
|
// parser would otherwise glue it onto the first field.
|
|
//
|
|
// The records stream: each one is parsed and its values are copied out
|
|
// as it arrives, so no row of text survives the row that produced it
|
|
// and the whole file is never held as strings. A malformed record
|
|
// stops the read and is reported first, then the first row whose field
|
|
// count differs from the first row's, then the first value that is not
|
|
// a number, the order a whole-file parse reports them in.
|
|
//
|
|
// The record parser is the hand-rolled tokenizer below, tuned for the
|
|
// numeric tables this entry point serves: it keeps the record
|
|
// semantics of encoding/csv for this caller's configuration, quotes
|
|
// and error reports included, while skipping the per-record string the
|
|
// standard parser builds.
|
|
func LoadCSVReader(r io.Reader, skipHeader bool) (*core.Array, error) {
|
|
br := bufio.NewReaderSize(r, csvReadBuffer)
|
|
prefix, _ := br.Peek(3)
|
|
if len(prefix) == 3 && prefix[0] == 0xEF && prefix[1] == 0xBB && prefix[2] == 0xBF {
|
|
_, _, _ = br.ReadRune() // consume the mark
|
|
}
|
|
t := &csvTokenizer{br: br}
|
|
vals := make([]float64, 0)
|
|
rows, cols := 0, 0
|
|
header := skipHeader
|
|
var ragged, badValue error
|
|
for {
|
|
rec, err := t.nextRecord()
|
|
if err == io.EOF {
|
|
break
|
|
}
|
|
if err != nil {
|
|
return nil, base.Errf("LoadCSV: %w", err)
|
|
}
|
|
if header {
|
|
header = false
|
|
continue
|
|
}
|
|
if rows == 0 {
|
|
cols = len(rec)
|
|
} else if len(rec) != cols {
|
|
// The read carries on to the end even after a defect, so a
|
|
// later record's syntax error still outranks an earlier
|
|
// ragged row, as the whole-file parse decides it. The row
|
|
// counter keeps advancing for the same reason.
|
|
if ragged == nil {
|
|
ragged = base.Errf("LoadCSV: row %d has %d fields, want %d", rows, len(rec), cols)
|
|
}
|
|
rows++
|
|
continue
|
|
}
|
|
if badValue == nil {
|
|
// The buffer doubles as the file arrives: append alone grows a
|
|
// large float slice by a quarter and rewrites it several times
|
|
// over, while doubling copies less than the final size once,
|
|
// and the first record sizes the buffer exactly.
|
|
if need := len(vals) + len(rec); need > cap(vals) {
|
|
grown := make([]float64, len(vals), max(2*cap(vals), need))
|
|
copy(grown, vals)
|
|
vals = grown
|
|
}
|
|
for j, f := range rec {
|
|
v, perr := strconv.ParseFloat(csvFieldText(f), 64)
|
|
if perr != nil {
|
|
if ne, ok := perr.(*strconv.NumError); ok {
|
|
// The message carries the field text: rebuild
|
|
// it over a private copy so it does not hang
|
|
// off the row buffer this read goes on
|
|
// overwriting.
|
|
perr = &strconv.NumError{Func: ne.Func, Num: string(f), Err: ne.Err}
|
|
}
|
|
badValue = base.Errf("LoadCSV: row %d col %d: %w", rows, j, perr)
|
|
break
|
|
}
|
|
vals = append(vals, v)
|
|
}
|
|
}
|
|
rows++
|
|
}
|
|
if ragged != nil {
|
|
return nil, ragged
|
|
}
|
|
if badValue != nil {
|
|
return nil, badValue
|
|
}
|
|
return core.FromFloats(vals, rows, cols)
|
|
}
|
|
|
|
// csvReadBuffer sizes the read buffer. A numeric table's rows run to a
|
|
// few hundred bytes, so one refill covers thousands of them and the
|
|
// reads stop being part of the cost.
|
|
const csvReadBuffer = 1 << 16
|
|
|
|
// csvFieldText views a field's bytes as the string strconv.ParseFloat
|
|
// parses, with no copy.
|
|
//
|
|
// SAFETY: the bytes live in the tokenizer's line or record buffer and
|
|
// stay untouched until the next record is read: ParseFloat reads the
|
|
// string only within the call, and on failure the NumError is rebuilt
|
|
// over a private copy before the loop can move on to the next record.
|
|
// On success no reference to the string escapes.
|
|
func csvFieldText(f []byte) string {
|
|
return unsafe.String(unsafe.SliceData(f), len(f))
|
|
}
|
|
|
|
// csvTokenizer is the record reader LoadCSVReader runs. It reproduces
|
|
// what encoding/csv's Reader does for the configuration this package
|
|
// reads with (comma separator, no comment character, no lazy quotes,
|
|
// no leading-space trimming, a variable field count): blank lines are
|
|
// skipped, a quote opens a quoted field, "" escapes one quote, \r\n is
|
|
// normalised to \n everywhere, interior newlines of a quoted field
|
|
// included, the last record may lack its newline, a trailing \r is
|
|
// dropped before EOF, and a malformed record comes back as a
|
|
// *csv.ParseError naming the same line and column the standard parser
|
|
// names. What it does not do is build the per-record string: the
|
|
// fields are byte ranges over the tokenizer's own buffers, valid until
|
|
// the next call.
|
|
type csvTokenizer struct {
|
|
br *bufio.Reader
|
|
|
|
numLine int // the line the reader sits on, counted from one
|
|
|
|
lineBuf []byte // assembles a record longer than the read buffer
|
|
recordBuf []byte // the record's unescaped fields, back to back
|
|
fields [][]byte // the current record's fields, reused
|
|
}
|
|
|
|
// lengthNL reports the number of bytes for the trailing \n.
|
|
func lengthNL(b []byte) int {
|
|
if len(b) > 0 && b[len(b)-1] == '\n' {
|
|
return 1
|
|
}
|
|
return 0
|
|
}
|
|
|
|
// readLine returns the next line with its end mark. A trailing \r\n is
|
|
// normalised to \n in place, a trailing \r is dropped before EOF, and
|
|
// every line read counts toward the line numbers the parse errors
|
|
// carry. The result is only valid until the next call.
|
|
func (t *csvTokenizer) readLine() ([]byte, error) {
|
|
line, err := t.br.ReadSlice('\n')
|
|
if err == bufio.ErrBufferFull {
|
|
t.lineBuf = append(t.lineBuf[:0], line...)
|
|
for err == bufio.ErrBufferFull {
|
|
line, err = t.br.ReadSlice('\n')
|
|
t.lineBuf = append(t.lineBuf, line...)
|
|
}
|
|
line = t.lineBuf
|
|
}
|
|
readSize := len(line)
|
|
if readSize > 0 && err == io.EOF {
|
|
err = nil
|
|
// For backwards compatibility, drop a trailing \r before EOF.
|
|
if line[readSize-1] == '\r' {
|
|
line = line[:readSize-1]
|
|
}
|
|
}
|
|
t.numLine++
|
|
// Normalise \r\n to \n on every input line.
|
|
if n := len(line); n >= 2 && line[n-2] == '\r' && line[n-1] == '\n' {
|
|
line[n-2] = '\n'
|
|
line = line[:n-1]
|
|
}
|
|
return line, err
|
|
}
|
|
|
|
// nextRecord parses the next record. The returned slices share the
|
|
// tokenizer's buffers and are only valid until the next call.
|
|
func (t *csvTokenizer) nextRecord() ([][]byte, error) {
|
|
// Read the record's first line, skipping the blank ones.
|
|
var line []byte
|
|
var errRead error
|
|
for errRead == nil {
|
|
line, errRead = t.readLine()
|
|
if errRead == nil && len(line) == lengthNL(line) {
|
|
continue // a blank line carries no record
|
|
}
|
|
break
|
|
}
|
|
if errRead == io.EOF {
|
|
return nil, errRead
|
|
}
|
|
|
|
// Fast path: a record with no quote anywhere splits on the commas
|
|
// of its one line. A quote is the only construct that carries a
|
|
// record across lines, escapes a field or raises a parse error,
|
|
// so the quoted path below owns every case beyond the split.
|
|
if bytes.IndexByte(line, '"') < 0 {
|
|
t.fields = t.fields[:0]
|
|
for {
|
|
i := bytes.IndexByte(line, ',')
|
|
if i < 0 {
|
|
t.fields = append(t.fields, line[:len(line)-lengthNL(line)])
|
|
return t.fields, errRead
|
|
}
|
|
t.fields = append(t.fields, line[:i])
|
|
line = line[i+1:]
|
|
}
|
|
}
|
|
|
|
recLine := t.numLine // the line the record starts on
|
|
t.recordBuf = t.recordBuf[:0]
|
|
t.fields = t.fields[:0]
|
|
posLine, posCol := t.numLine, 1
|
|
var parseErr error
|
|
parseField:
|
|
for {
|
|
if len(line) == 0 || line[0] != '"' {
|
|
// Non-quoted field: everything up to the comma, with
|
|
// the end mark stripped.
|
|
i := bytes.IndexByte(line, ',')
|
|
field := line
|
|
if i >= 0 {
|
|
field = field[:i]
|
|
} else {
|
|
field = field[:len(field)-lengthNL(field)]
|
|
}
|
|
// A quote may not appear in a non-quoted field.
|
|
if j := bytes.IndexByte(field, '"'); j >= 0 {
|
|
parseErr = &csv.ParseError{StartLine: recLine, Line: t.numLine, Column: posCol + j, Err: csv.ErrBareQuote}
|
|
break parseField
|
|
}
|
|
start := len(t.recordBuf)
|
|
t.recordBuf = append(t.recordBuf, field...)
|
|
t.fields = append(t.fields, t.recordBuf[start:len(t.recordBuf)])
|
|
if i >= 0 {
|
|
line = line[i+1:]
|
|
posCol += i + 1
|
|
continue parseField
|
|
}
|
|
break parseField
|
|
}
|
|
// Quoted field: the opening quote is consumed, the rest
|
|
// accumulates until the closing one.
|
|
fieldStart := len(t.recordBuf)
|
|
line = line[1:]
|
|
posCol++
|
|
for {
|
|
i := bytes.IndexByte(line, '"')
|
|
if i >= 0 {
|
|
t.recordBuf = append(t.recordBuf, line[:i]...)
|
|
line = line[i+1:]
|
|
posCol += i + 1
|
|
switch {
|
|
case len(line) > 0 && line[0] == '"':
|
|
// "" escapes one quote.
|
|
t.recordBuf = append(t.recordBuf, '"')
|
|
line = line[1:]
|
|
posCol++
|
|
case len(line) > 0 && line[0] == ',':
|
|
// ", closes the field.
|
|
line = line[1:]
|
|
posCol++
|
|
t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)])
|
|
continue parseField
|
|
case len(line) == 0 || (len(line) == 1 && line[0] == '\n'):
|
|
// A closing quote at the end of the line closes
|
|
// the field and the record with it.
|
|
t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)])
|
|
break parseField
|
|
default:
|
|
// Anything after a closing quote that is neither
|
|
// comma nor end of line.
|
|
parseErr = &csv.ParseError{StartLine: recLine, Line: t.numLine, Column: posCol - 1, Err: csv.ErrQuote}
|
|
break parseField
|
|
}
|
|
} else if len(line) > 0 {
|
|
// End of line inside the field: the whole line, end
|
|
// mark included, belongs to the field.
|
|
t.recordBuf = append(t.recordBuf, line...)
|
|
if errRead != nil {
|
|
break parseField
|
|
}
|
|
posCol += len(line)
|
|
line, errRead = t.readLine()
|
|
if len(line) > 0 {
|
|
posLine++
|
|
posCol = 1
|
|
}
|
|
if errRead == io.EOF {
|
|
errRead = nil
|
|
}
|
|
} else {
|
|
// End of input inside the field.
|
|
if errRead == nil {
|
|
parseErr = &csv.ParseError{StartLine: recLine, Line: posLine, Column: posCol, Err: csv.ErrQuote}
|
|
break parseField
|
|
}
|
|
t.fields = append(t.fields, t.recordBuf[fieldStart:len(t.recordBuf)])
|
|
break parseField
|
|
}
|
|
}
|
|
}
|
|
if parseErr == nil {
|
|
parseErr = errRead
|
|
}
|
|
return t.fields, parseErr
|
|
}
|