summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-03 19:20:57 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-03 20:06:42 +0200
commitfda1267d7d6600c04a6cd7eb2a818ebefb8999e4 (patch)
treecbfab8e35e396a8b5dc4e9d5ed765ad699dac1db
parentc1644da32ecb9856f15e2c919a84d2666d5ef29c (diff)
Return document index in text iterator
-rw-r--r--dataset/cmd/dataset/main.go8
-rw-r--r--dataset/parquet.go122
-rw-r--r--dataset/reader.go2
-rw-r--r--dataset/string.go6
-rw-r--r--llm/cmd/eval/ppl.go2
-rw-r--r--llm/run.go23
-rw-r--r--tokenizer/bpe/cmd/tokenize/main.go6
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 != "" {