From 236d184ad8a2af0e6416b381ea55668c89d40df2 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Mon, 16 Mar 2026 17:15:24 +0100 Subject: Rename dataset reader --- dataset/cmd/dataset/main.go | 4 +- dataset/parquet.go | 135 ++++++++++++++++++++++++++++++++++++++++++++ dataset/reader.go | 135 -------------------------------------------- llm/cmd/ppl/ppl.go | 6 +- llm/perplexity.go | 2 +- 5 files changed, 141 insertions(+), 141 deletions(-) create mode 100644 dataset/parquet.go delete mode 100644 dataset/reader.go diff --git a/dataset/cmd/dataset/main.go b/dataset/cmd/dataset/main.go index 5fd633b..8bd892d 100644 --- a/dataset/cmd/dataset/main.go +++ b/dataset/cmd/dataset/main.go @@ -19,9 +19,9 @@ func main() { split := filepath.Join(root, "train") - var reader *dataset.Reader + var reader *dataset.ParquetReader - if r, err := dataset.NewReader(split); err != nil { + if r, err := dataset.NewParquetReader(split); err != nil { log.Fatal(err) } else { reader = r diff --git a/dataset/parquet.go b/dataset/parquet.go new file mode 100644 index 0000000..c80758f --- /dev/null +++ b/dataset/parquet.go @@ -0,0 +1,135 @@ +package dataset + +import ( + "context" + "errors" + "fmt" + "iter" + "path/filepath" + "slices" + + "github.com/apache/arrow-go/v18/arrow" + "github.com/apache/arrow-go/v18/arrow/array" + "github.com/apache/arrow-go/v18/arrow/memory" + "github.com/apache/arrow-go/v18/parquet/file" + "github.com/apache/arrow-go/v18/parquet/pqarrow" +) + +type ParquetReader struct { + shards []string + batchSize int64 + err error +} + +func NewParquetReader(name string) (*ParquetReader, error) { + var shards []string + + if matches, err := filepath.Glob(filepath.Join(name, "*.parquet")); err != nil { + return nil, err + } else { + shards = matches + } + + if len(shards) == 0 { + return nil, fmt.Errorf("no parquet files in %q", name) + } + + slices.Sort(shards) + + r := &ParquetReader{ + shards: shards, + batchSize: 1024, + } + + return r, nil +} + +func (r *ParquetReader) Err() error { + return r.err +} + +func (r *ParquetReader) Texts(column string) iter.Seq[string] { + return func(yield func(string) bool) { + for _, name := range r.shards { + err := read(name, column, r.batchSize, yield) + + if errors.Is(err, stop) { + return + } + + if err != nil { + r.err = err + + return + } + } + } +} + +var stop = errors.New("stop") + +func read(name, column string, batchSize int64, yield func(string) bool) error { + var reader *file.Reader + + if r, err := file.OpenParquetFile(name, false); err != nil { + return err + } else { + reader = r + + defer reader.Close() + } + + var arrowReader *pqarrow.FileReader + + if r, err := pqarrow.NewFileReader(reader, pqarrow.ArrowReadProperties{BatchSize: batchSize}, memory.DefaultAllocator); err != nil { + return err + } else { + arrowReader = r + } + + var schema *arrow.Schema + + if s, err := arrowReader.Schema(); err != nil { + return err + } else { + schema = s + } + + idxs := schema.FieldIndices(column) + + if len(idxs) == 0 { + return errors.New("unknown column") + } + + var recordReader pqarrow.RecordReader + + if r, err := arrowReader.GetRecordReader(context.TODO(), []int{idxs[0]}, nil); err != nil { + return err + } else { + recordReader = r + + defer recordReader.Release() + } + + for recordReader.Next() { + record := recordReader.RecordBatch() + + text, ok := record.Column(0).(*array.String) + + if !ok { + return fmt.Errorf("unexpected column type") + } + + for i := range int(record.NumRows()) { + if !text.IsValid(i) { + continue + } + + if !yield(text.Value(i)) { + return stop + } + } + } + + return nil +} diff --git a/dataset/reader.go b/dataset/reader.go deleted file mode 100644 index e882332..0000000 --- a/dataset/reader.go +++ /dev/null @@ -1,135 +0,0 @@ -package dataset - -import ( - "context" - "errors" - "fmt" - "iter" - "path/filepath" - "slices" - - "github.com/apache/arrow-go/v18/arrow" - "github.com/apache/arrow-go/v18/arrow/array" - "github.com/apache/arrow-go/v18/arrow/memory" - "github.com/apache/arrow-go/v18/parquet/file" - "github.com/apache/arrow-go/v18/parquet/pqarrow" -) - -type Reader struct { - shards []string - batchSize int64 - err error -} - -func NewReader(name string) (*Reader, error) { - var shards []string - - if matches, err := filepath.Glob(filepath.Join(name, "*.parquet")); err != nil { - return nil, err - } else { - shards = matches - } - - if len(shards) == 0 { - return nil, fmt.Errorf("no parquet files in %q", name) - } - - slices.Sort(shards) - - r := &Reader{ - shards: shards, - batchSize: 1024, - } - - return r, nil -} - -func (r *Reader) Err() error { - return r.err -} - -func (r *Reader) Texts(column string) iter.Seq[string] { - return func(yield func(string) bool) { - for _, name := range r.shards { - err := read(name, column, r.batchSize, yield) - - if errors.Is(err, stop) { - return - } - - if err != nil { - r.err = err - - return - } - } - } -} - -var stop = errors.New("stop") - -func read(name, column string, batchSize int64, yield func(string) bool) error { - var reader *file.Reader - - if r, err := file.OpenParquetFile(name, false); err != nil { - return err - } else { - reader = r - - defer reader.Close() - } - - var arrowReader *pqarrow.FileReader - - if r, err := pqarrow.NewFileReader(reader, pqarrow.ArrowReadProperties{BatchSize: batchSize}, memory.DefaultAllocator); err != nil { - return err - } else { - arrowReader = r - } - - var schema *arrow.Schema - - if s, err := arrowReader.Schema(); err != nil { - return err - } else { - schema = s - } - - idxs := schema.FieldIndices(column) - - if len(idxs) == 0 { - return errors.New("unknown column") - } - - var recordReader pqarrow.RecordReader - - if r, err := arrowReader.GetRecordReader(context.TODO(), []int{idxs[0]}, nil); err != nil { - return err - } else { - recordReader = r - - defer recordReader.Release() - } - - for recordReader.Next() { - record := recordReader.RecordBatch() - - text, ok := record.Column(0).(*array.String) - - if !ok { - return fmt.Errorf("unexpected column type") - } - - for i := range int(record.NumRows()) { - if !text.IsValid(i) { - continue - } - - if !yield(text.Value(i)) { - return stop - } - } - } - - return nil -} diff --git a/llm/cmd/ppl/ppl.go b/llm/cmd/ppl/ppl.go index 584cd16..6912e48 100644 --- a/llm/cmd/ppl/ppl.go +++ b/llm/cmd/ppl/ppl.go @@ -30,10 +30,10 @@ func main() { fmt.Println(ppl) } -func data() *dataset.Reader { - var miniPile *dataset.Reader +func data() *dataset.ParquetReader { + var miniPile *dataset.ParquetReader - if r, err := dataset.NewReader("dataset/cmd/dataset/tmp/minipile/validation"); err != nil { + if r, err := dataset.NewParquetReader("dataset/cmd/dataset/tmp/minipile/validation"); err != nil { log.Fatal(err) } else { miniPile = r diff --git a/llm/perplexity.go b/llm/perplexity.go index 9596e8d..2240f2d 100644 --- a/llm/perplexity.go +++ b/llm/perplexity.go @@ -14,7 +14,7 @@ import ( "go.jknobloc.com/x/dataset" ) -func (e *Evaluator) Perplexity(data *dataset.Reader, window, stride int) (float64, error) { +func (e *Evaluator) Perplexity(data *dataset.ParquetReader, window, stride int) (float64, error) { docs := make([]string, 0) n := 0 -- cgit v1.3.1