summaryrefslogtreecommitdiff
path: root/dataset
diff options
context:
space:
mode:
Diffstat (limited to 'dataset')
-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
4 files changed, 73 insertions, 65 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
}
}