diff options
| -rw-r--r-- | dataset/cmd/dataset/main.go | 4 | ||||
| -rw-r--r-- | dataset/parquet.go (renamed from dataset/reader.go) | 10 | ||||
| -rw-r--r-- | llm/cmd/ppl/ppl.go | 6 | ||||
| -rw-r--r-- | llm/perplexity.go | 2 |
4 files changed, 11 insertions, 11 deletions
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/reader.go b/dataset/parquet.go index e882332..c80758f 100644 --- a/dataset/reader.go +++ b/dataset/parquet.go @@ -15,13 +15,13 @@ import ( "github.com/apache/arrow-go/v18/parquet/pqarrow" ) -type Reader struct { +type ParquetReader struct { shards []string batchSize int64 err error } -func NewReader(name string) (*Reader, error) { +func NewParquetReader(name string) (*ParquetReader, error) { var shards []string if matches, err := filepath.Glob(filepath.Join(name, "*.parquet")); err != nil { @@ -36,7 +36,7 @@ func NewReader(name string) (*Reader, error) { slices.Sort(shards) - r := &Reader{ + r := &ParquetReader{ shards: shards, batchSize: 1024, } @@ -44,11 +44,11 @@ func NewReader(name string) (*Reader, error) { return r, nil } -func (r *Reader) Err() error { +func (r *ParquetReader) Err() error { return r.err } -func (r *Reader) Texts(column string) iter.Seq[string] { +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) 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 |
