Files
tensor/io/csv.go
T

472 lines
15 KiB
Go
Raw Normal View History

2026-09-03 10:00:00 +02:00
// 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
}