summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-03 16:02:05 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-03 18:45:57 +0200
commitc1644da32ecb9856f15e2c919a84d2666d5ef29c (patch)
tree0c19b78e35c05d7814f855128642a384bd8660cd
parenta61bc9ff890ea50d41a406cd1ffe554341551d2b (diff)
Scope text column to parquet reader
-rw-r--r--dataset/cmd/dataset/main.go2
-rw-r--r--dataset/parquet.go16
-rw-r--r--dataset/reader.go2
-rw-r--r--dataset/string.go2
-rw-r--r--llm/cmd/eval/ppl.go2
-rw-r--r--llm/run.go2
-rw-r--r--tokenizer/bpe/cmd/tokenize/main.go2
7 files changed, 15 insertions, 13 deletions
diff --git a/dataset/cmd/dataset/main.go b/dataset/cmd/dataset/main.go
index 8bd892d..c89ffc9 100644
--- a/dataset/cmd/dataset/main.go
+++ b/dataset/cmd/dataset/main.go
@@ -29,7 +29,7 @@ func main() {
i := 0
- for s := range reader.Texts("text") {
+ for s := range reader.Texts() {
_ = s
fmt.Printf("\r%d", i+1)
i++
diff --git a/dataset/parquet.go b/dataset/parquet.go
index c80758f..8b91458 100644
--- a/dataset/parquet.go
+++ b/dataset/parquet.go
@@ -16,9 +16,10 @@ import (
)
type ParquetReader struct {
- shards []string
- batchSize int64
- err error
+ shards []string
+ textColumn string
+ batchSize int64
+ err error
}
func NewParquetReader(name string) (*ParquetReader, error) {
@@ -37,8 +38,9 @@ func NewParquetReader(name string) (*ParquetReader, error) {
slices.Sort(shards)
r := &ParquetReader{
- shards: shards,
- batchSize: 1024,
+ shards: shards,
+ textColumn: "text",
+ batchSize: 1024,
}
return r, nil
@@ -48,10 +50,10 @@ func (r *ParquetReader) Err() error {
return r.err
}
-func (r *ParquetReader) Texts(column string) iter.Seq[string] {
+func (r *ParquetReader) Texts() iter.Seq[string] {
return func(yield func(string) bool) {
for _, name := range r.shards {
- err := read(name, column, r.batchSize, yield)
+ err := read(name, r.textColumn, r.batchSize, yield)
if errors.Is(err, stop) {
return
diff --git a/dataset/reader.go b/dataset/reader.go
index 9388ba5..a80d401 100644
--- a/dataset/reader.go
+++ b/dataset/reader.go
@@ -3,6 +3,6 @@ package dataset
import "iter"
type Reader interface {
- Texts(column string) iter.Seq[string]
+ Texts() iter.Seq[string]
Err() error
}
diff --git a/dataset/string.go b/dataset/string.go
index 4c1112c..f98af96 100644
--- a/dataset/string.go
+++ b/dataset/string.go
@@ -12,7 +12,7 @@ func NewStringReader(s string) *StringReader {
}
}
-func (s *StringReader) Texts(_ string) iter.Seq[string] {
+func (s *StringReader) Texts() iter.Seq[string] {
return func(yield func(string) bool) {
if !yield(s.s) {
return
diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go
index 60dd30a..679de90 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("text") {
+ for d := range miniPile.Texts() {
docs = append(docs, d)
}
diff --git a/llm/run.go b/llm/run.go
index 1843c94..483cd5d 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -63,7 +63,7 @@ func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int
b := newBatch(e.batchSize)
- for d := range data.Texts("text") {
+ for d := range data.Texts() {
for w, s := range tb.Push(n, d) {
if s == 0 {
s = 1 // first token as context
diff --git a/tokenizer/bpe/cmd/tokenize/main.go b/tokenizer/bpe/cmd/tokenize/main.go
index 1dbfaae..a0519cf 100644
--- a/tokenizer/bpe/cmd/tokenize/main.go
+++ b/tokenizer/bpe/cmd/tokenize/main.go
@@ -55,7 +55,7 @@ func main() {
n := 0
- for d := range reader.Texts("text") {
+ for d := range reader.Texts() {
if n >= 1000 {
break
}