diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-03 19:20:57 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-03 20:06:42 +0200 |
| commit | fda1267d7d6600c04a6cd7eb2a818ebefb8999e4 (patch) | |
| tree | cbfab8e35e396a8b5dc4e9d5ed765ad699dac1db | |
| parent | c1644da32ecb9856f15e2c919a84d2666d5ef29c (diff) | |
Return document index in text iterator
| -rw-r--r-- | dataset/cmd/dataset/main.go | 8 | ||||
| -rw-r--r-- | dataset/parquet.go | 122 | ||||
| -rw-r--r-- | dataset/reader.go | 2 | ||||
| -rw-r--r-- | dataset/string.go | 6 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 2 | ||||
| -rw-r--r-- | llm/run.go | 23 | ||||
| -rw-r--r-- | 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) } @@ -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 != "" { |
