From fda1267d7d6600c04a6cd7eb2a818ebefb8999e4 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Fri, 3 Apr 2026 19:20:57 +0200 Subject: Return document index in text iterator --- dataset/cmd/dataset/main.go | 8 +-- dataset/parquet.go | 122 ++++++++++++++++++++----------------- dataset/reader.go | 2 +- dataset/string.go | 6 +- llm/cmd/eval/ppl.go | 2 +- llm/run.go | 23 +++---- tokenizer/bpe/cmd/tokenize/main.go | 6 +- 7 files changed, 87 insertions(+), 82 deletions(-) diff --git a/dataset/cmd/dataset/main.go b/dataset/cmd/dataset/main.go index c89ffc9..66d2f1b 100644 --- a/dataset/cmd/dataset/main.go +++ b/dataset/cmd/dataset/main.go @@ -27,12 +27,8 @@ func main() { reader = r } - i := 0 - - for s := range reader.Texts() { - _ = s - fmt.Printf("\r%d", i+1) - i++ + for n := range reader.Texts() { + fmt.Printf("\r%d", n) } if err := reader.Err(); err != nil { diff --git a/dataset/parquet.go b/dataset/parquet.go index 8b91458..f08c332 100644 --- a/dataset/parquet.go +++ b/dataset/parquet.go @@ -46,92 +46,104 @@ func NewParquetReader(name string) (*ParquetReader, error) { return r, nil } -func (r *ParquetReader) Err() error { - return r.err +func (p *ParquetReader) Err() error { + return p.err } -func (r *ParquetReader) Texts() iter.Seq[string] { - return func(yield func(string) bool) { - for _, name := range r.shards { - err := read(name, r.textColumn, r.batchSize, yield) +func (p *ParquetReader) Texts() iter.Seq2[int, string] { + return func(yield func(int, string) bool) { + n := 0 - if errors.Is(err, stop) { - return - } + for _, name := range p.shards { + for text := range p.read(name) { + if p.err != nil { + return + } - if err != nil { - r.err = err + if !yield(n, text) { + return + } - return + n++ } } } } -var stop = errors.New("stop") +func (p *ParquetReader) read(name string) iter.Seq[string] { + return func(yield func(string) bool) { + var reader *file.Reader -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 { + p.err = err - if r, err := file.OpenParquetFile(name, false); err != nil { - return err - } else { - reader = r + return + } else { + reader = r + } defer reader.Close() - } - var arrowReader *pqarrow.FileReader + var arrowReader *pqarrow.FileReader - if r, err := pqarrow.NewFileReader(reader, pqarrow.ArrowReadProperties{BatchSize: batchSize}, memory.DefaultAllocator); err != nil { - return err - } else { - arrowReader = r - } + if r, err := pqarrow.NewFileReader(reader, pqarrow.ArrowReadProperties{BatchSize: p.batchSize}, memory.DefaultAllocator); err != nil { + p.err = err - var schema *arrow.Schema + return + } else { + arrowReader = r + } - if s, err := arrowReader.Schema(); err != nil { - return err - } else { - schema = s - } + var schema *arrow.Schema - idxs := schema.FieldIndices(column) + if s, err := arrowReader.Schema(); err != nil { + p.err = err - if len(idxs) == 0 { - return errors.New("unknown column") - } + return + } else { + schema = s + } - var recordReader pqarrow.RecordReader + idxs := schema.FieldIndices(p.textColumn) - if r, err := arrowReader.GetRecordReader(context.TODO(), []int{idxs[0]}, nil); err != nil { - return err - } else { - recordReader = r + if len(idxs) == 0 { + p.err = errors.New("unknown column") - defer recordReader.Release() - } + return + } - for recordReader.Next() { - record := recordReader.RecordBatch() + var recordReader pqarrow.RecordReader - text, ok := record.Column(0).(*array.String) + if r, err := arrowReader.GetRecordReader(context.TODO(), []int{idxs[0]}, nil); err != nil { + p.err = err - if !ok { - return fmt.Errorf("unexpected column type") + return + } else { + recordReader = r } - for i := range int(record.NumRows()) { - if !text.IsValid(i) { - continue + defer recordReader.Release() + + for recordReader.Next() { + record := recordReader.RecordBatch() + + text, ok := record.Column(0).(*array.String) + + if !ok { + p.err = fmt.Errorf("unexpected column type") + + return } - if !yield(text.Value(i)) { - return stop + for i := range int(record.NumRows()) { + if !text.IsValid(i) { + continue + } + + if !yield(text.Value(i)) { + return + } } } } - - return nil } diff --git a/dataset/reader.go b/dataset/reader.go index a80d401..6458010 100644 --- a/dataset/reader.go +++ b/dataset/reader.go @@ -3,6 +3,6 @@ package dataset import "iter" type Reader interface { - Texts() iter.Seq[string] + Texts() iter.Seq2[int, string] Err() error } diff --git a/dataset/string.go b/dataset/string.go index f98af96..c1f7600 100644 --- a/dataset/string.go +++ b/dataset/string.go @@ -12,9 +12,9 @@ func NewStringReader(s string) *StringReader { } } -func (s *StringReader) Texts() iter.Seq[string] { - return func(yield func(string) bool) { - if !yield(s.s) { +func (s *StringReader) Texts() iter.Seq2[int, string] { + return func(yield func(int, string) bool) { + if !yield(0, s.s) { return } } diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go index 679de90..261114e 100644 --- a/llm/cmd/eval/ppl.go +++ b/llm/cmd/eval/ppl.go @@ -53,7 +53,7 @@ func joined() dataset.Reader { docs := make([]string, 0) - for d := range miniPile.Texts() { + for _, d := range miniPile.Texts() { docs = append(docs, d) } diff --git a/llm/run.go b/llm/run.go index 483cd5d..fd77ce9 100644 --- a/llm/run.go +++ b/llm/run.go @@ -52,9 +52,6 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int return int(e.completed.Load()) }) - n := 0 - m := 0 - defer pb.Close() tb := NewTokenBuffer(e.tokenizer, window, stride) @@ -63,20 +60,28 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int b := newBatch(e.batchSize) - for d := range data.Texts() { + doc := 0 + pos := 0 + + for n, d := range data.Texts() { + if n != doc { + pos = 0 + doc = n + } + for w, s := range tb.Push(n, d) { if s == 0 { s = 1 // first token as context } b.AddJob(Job{ - Document: n, - Position: m, + Document: doc, + Position: pos, Tokens: w, Seen: s, }) - m++ + pos++ if b.Size() == e.batchSize { e.jobs <- *b @@ -88,10 +93,6 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int pb.SetTotal(int(e.scheduled.Load())) } } - - n++ - - m = 0 } if s := b.Size(); s > 0 { diff --git a/tokenizer/bpe/cmd/tokenize/main.go b/tokenizer/bpe/cmd/tokenize/main.go index a0519cf..08d88aa 100644 --- a/tokenizer/bpe/cmd/tokenize/main.go +++ b/tokenizer/bpe/cmd/tokenize/main.go @@ -53,9 +53,7 @@ func main() { defer pb.Close() - n := 0 - - for d := range reader.Texts() { + for n, d := range reader.Texts() { if n >= 1000 { break } @@ -65,8 +63,6 @@ func main() { _ = tokens processed.Add(1) - - n++ } if *memprofile != "" { -- cgit v1.2.3