summaryrefslogtreecommitdiff
path: root/dataset
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-02-28 18:11:20 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-03-02 17:30:02 +0100
commit56a5c2d209192901bde81a51128c79f222f5c5be (patch)
treef6281156f41f88328e59a8f61429f1d874a46ac7 /dataset
parent8103b598b73f2bac8ed3e72821998a7e4881f832 (diff)
Add dataset package
Diffstat (limited to 'dataset')
-rw-r--r--dataset/cmd/dataset/main.go43
-rw-r--r--dataset/download.go224
-rw-r--r--dataset/go.mod34
-rw-r--r--dataset/go.sum90
-rw-r--r--dataset/reader.go135
5 files changed, 526 insertions, 0 deletions
diff --git a/dataset/cmd/dataset/main.go b/dataset/cmd/dataset/main.go
new file mode 100644
index 0000000..76028be
--- /dev/null
+++ b/dataset/cmd/dataset/main.go
@@ -0,0 +1,43 @@
+package main
+
+import (
+ "fmt"
+ "log"
+ "path/filepath"
+
+ "github.com/jonasknobloch/x/dataset"
+)
+
+func main() {
+ root := "./tmp/minipile"
+
+ if err := dataset.Download("JeanKaddour/minipile", "", root); err != nil {
+ log.Fatal(err)
+ }
+
+ fmt.Println()
+
+ split := filepath.Join(root, "train")
+
+ var reader *dataset.Reader
+
+ if r, err := dataset.NewReader(split); err != nil {
+ log.Fatal(err)
+ } else {
+ reader = r
+ }
+
+ i := 0
+
+ for s := range reader.Texts("text") {
+ _ = s
+ fmt.Printf("\r%d", i+1)
+ i++
+ }
+
+ if err := reader.Err(); err != nil {
+ log.Fatal(err)
+ }
+
+ fmt.Println()
+}
diff --git a/dataset/download.go b/dataset/download.go
new file mode 100644
index 0000000..0df6ebe
--- /dev/null
+++ b/dataset/download.go
@@ -0,0 +1,224 @@
+package dataset
+
+import (
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "iter"
+ "net/http"
+ "net/url"
+ "os"
+ "path/filepath"
+ "slices"
+)
+
+const Base = "https://datasets-server.huggingface.co/parquet"
+
+func Download(dataset, config, root string) error {
+ var pl *parquetListing
+
+ if l, err := listing(dataset, config); err != nil {
+ return err
+ } else {
+ pl = l
+ }
+
+ splits := pl.Splits()
+
+ for _, split := range splits {
+ dir := filepath.Join(root, split)
+
+ if err := os.MkdirAll(dir, 0o755); err != nil {
+ return err
+ }
+
+ total := int64(0)
+
+ for pf := range pl.Split(split) {
+ name := filepath.Join(dir, filepath.Base(pf.Filename))
+
+ if n, err := download(pf, name); err != nil {
+ return err
+ } else {
+ total += n
+
+ fmt.Printf("[%s] downloaded %s (%d bytes), total=%d\n", split, filepath.Base(name), n, total)
+ }
+ }
+ }
+
+ return nil
+}
+
+type parquetListing struct {
+ ParquetFiles []parquetFile `json:"parquet_files"`
+}
+
+type parquetFile struct {
+ Dataset string `json:"dataset"`
+ Config string `json:"config"`
+ Split string `json:"split"`
+ URL string `json:"url"`
+ Filename string `json:"filename"`
+ Size int64 `json:"size"`
+}
+
+func (pl *parquetListing) Splits() []string {
+ seen := make(map[string]struct{})
+
+ for _, ps := range pl.ParquetFiles {
+ if ps.Split == "" {
+ continue // TODO possible?
+ }
+
+ seen[ps.Split] = struct{}{}
+ }
+
+ splits := make([]string, 0)
+
+ for k := range seen {
+ splits = append(splits, k)
+ }
+
+ slices.Sort(splits)
+
+ return splits
+}
+
+func (pl *parquetListing) Split(split string) iter.Seq[parquetFile] {
+ return func(yield func(pf parquetFile) bool) {
+ for _, pf := range pl.ParquetFiles {
+ if pf.Split != split {
+ continue
+ }
+
+ if !yield(pf) {
+ return
+ }
+ }
+ }
+}
+
+func listing(dataset, config string) (*parquetListing, error) {
+ var parquet *url.URL
+
+ if u, err := url.Parse(Base); err != nil {
+ return nil, err
+ } else {
+ parquet = u
+ }
+
+ query := parquet.Query()
+
+ query.Set("dataset", dataset)
+ query.Set("config", config)
+
+ parquet.RawQuery = query.Encode()
+
+ var req *http.Request
+
+ if r, err := http.NewRequest("GET", parquet.String(), nil); err != nil {
+ return nil, err
+ } else {
+ req = r
+ }
+
+ var resp *http.Response
+
+ if r, err := http.DefaultClient.Do(req); err != nil {
+ return nil, err
+ } else {
+ resp = r
+
+ defer resp.Body.Close()
+ }
+
+ if resp.StatusCode != http.StatusOK {
+ return nil, errors.New("http not ok")
+ }
+
+ var pl parquetListing
+
+ if err := json.NewDecoder(resp.Body).Decode(&pl); err != nil {
+ return nil, err
+ }
+
+ return &pl, nil
+}
+
+func download(pf parquetFile, name string) (int64, error) {
+ if st, err := os.Stat(name); err == nil && st.Size() == pf.Size {
+ return pf.Size, nil
+ }
+
+ var req *http.Request
+
+ if r, err := http.NewRequest("GET", pf.URL, nil); err != nil {
+ return 0, err
+ } else {
+ req = r
+ }
+
+ var resp *http.Response
+
+ if r, err := http.DefaultClient.Do(req); err != nil {
+ return 0, err
+ } else {
+ resp = r
+
+ defer resp.Body.Close()
+ }
+
+ if resp.StatusCode != http.StatusOK {
+ return 0, errors.New("http not ok")
+ }
+
+ tmp := name + ".part"
+
+ var file *os.File
+
+ ok := false
+
+ if f, err := os.Create(tmp); err != nil {
+ return 0, err
+ } else {
+ file = f
+ }
+
+ defer func() {
+ _ = file.Close()
+
+ if !ok {
+ _ = os.Remove(tmp)
+ }
+ }()
+
+ written := int64(0)
+
+ if n, err := io.Copy(file, resp.Body); err != nil {
+ return 0, err
+ } else {
+ written = n
+ }
+
+ if err := file.Sync(); err != nil {
+ return 0, err
+ }
+
+ if pf.Size > 0 && written != pf.Size {
+ return 0, errors.New("size mismatch")
+ }
+
+ if err := file.Close(); err != nil {
+ return 0, err
+ }
+
+ if err := os.Rename(tmp, name); err != nil {
+ return 0, err
+ }
+
+ ok = true
+
+ return written, nil
+}
diff --git a/dataset/go.mod b/dataset/go.mod
new file mode 100644
index 0000000..6b8a6c6
--- /dev/null
+++ b/dataset/go.mod
@@ -0,0 +1,34 @@
+module github.com/jonasknobloch/x/dataset
+
+go 1.25.0
+
+require github.com/apache/arrow-go/v18 v18.5.1
+
+require (
+ github.com/andybalholm/brotli v1.2.0 // indirect
+ github.com/apache/thrift v0.22.0 // indirect
+ github.com/cespare/xxhash/v2 v2.3.0 // indirect
+ github.com/goccy/go-json v0.10.5 // indirect
+ github.com/golang/snappy v1.0.0 // indirect
+ github.com/google/flatbuffers v25.12.19+incompatible // indirect
+ github.com/google/uuid v1.6.0 // indirect
+ github.com/klauspost/asmfmt v1.3.2 // indirect
+ github.com/klauspost/compress v1.18.2 // indirect
+ github.com/klauspost/cpuid/v2 v2.3.0 // indirect
+ github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 // indirect
+ github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3 // indirect
+ github.com/pierrec/lz4/v4 v4.1.23 // indirect
+ github.com/zeebo/xxh3 v1.0.2 // indirect
+ golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 // indirect
+ golang.org/x/mod v0.32.0 // indirect
+ golang.org/x/net v0.49.0 // indirect
+ golang.org/x/sync v0.19.0 // indirect
+ golang.org/x/sys v0.40.0 // indirect
+ golang.org/x/telemetry v0.0.0-20260109210033-bd525da824e2 // indirect
+ golang.org/x/text v0.33.0 // indirect
+ golang.org/x/tools v0.41.0 // indirect
+ golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da // indirect
+ google.golang.org/genproto/googleapis/rpc v0.0.0-20251029180050-ab9386a59fda // indirect
+ google.golang.org/grpc v1.78.0 // indirect
+ google.golang.org/protobuf v1.36.11 // indirect
+)
diff --git a/dataset/go.sum b/dataset/go.sum
new file mode 100644
index 0000000..48e1447
--- /dev/null
+++ b/dataset/go.sum
@@ -0,0 +1,90 @@
+github.com/andybalholm/brotli v1.2.0 h1:ukwgCxwYrmACq68yiUqwIWnGY0cTPox/M94sVwToPjQ=
+github.com/andybalholm/brotli v1.2.0/go.mod h1:rzTDkvFWvIrjDXZHkuS16NPggd91W3kUSvPlQ1pLaKY=
+github.com/apache/arrow-go/v18 v18.5.1 h1:yaQ6zxMGgf9YCYw4/oaeOU3AULySDlAYDOcnr4LdHdI=
+github.com/apache/arrow-go/v18 v18.5.1/go.mod h1:OCCJsmdq8AsRm8FkBSSmYTwL/s4zHW9CqxeBxEytkNE=
+github.com/apache/thrift v0.22.0 h1:r7mTJdj51TMDe6RtcmNdQxgn9XcyfGDOzegMDRg47uc=
+github.com/apache/thrift v0.22.0/go.mod h1:1e7J/O1Ae6ZQMTYdy9xa3w9k+XHWPfRvdPyJeynQ+/g=
+github.com/cespare/xxhash/v2 v2.3.0 h1:UL815xU9SqsFlibzuggzjXhog7bL6oX9BbNZnL2UFvs=
+github.com/cespare/xxhash/v2 v2.3.0/go.mod h1:VGX0DQ3Q6kWi7AoAeZDth3/j3BFtOZR5XLFGgcrjCOs=
+github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc h1:U9qPSI2PIWSS1VwoXQT9A3Wy9MM3WgvqSxFWenqJduM=
+github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc/go.mod h1:J7Y8YcW2NihsgmVo/mv3lAwl/skON4iLHjSsI+c5H38=
+github.com/go-logr/logr v1.4.3 h1:CjnDlHq8ikf6E492q6eKboGOC0T8CDaOvkHCIg8idEI=
+github.com/go-logr/logr v1.4.3/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY=
+github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag=
+github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE=
+github.com/goccy/go-json v0.10.5 h1:Fq85nIqj+gXn/S5ahsiTlK3TmC85qgirsdTP/+DeaC4=
+github.com/goccy/go-json v0.10.5/go.mod h1:oq7eo15ShAhp70Anwd5lgX2pLfOS3QCiwU/PULtXL6M=
+github.com/golang/protobuf v1.5.4 h1:i7eJL8qZTpSEXOPTxNKhASYpMn+8e5Q6AdndVa1dWek=
+github.com/golang/protobuf v1.5.4/go.mod h1:lnTiLA8Wa4RWRcIUkrtSVa5nRhsEGBg48fD6rSs7xps=
+github.com/golang/snappy v1.0.0 h1:Oy607GVXHs7RtbggtPBnr2RmDArIsAefDwvrdWvRhGs=
+github.com/golang/snappy v1.0.0/go.mod h1:/XxbfmMg8lxefKM7IXC3fBNl/7bRcc72aCRzEWrmP2Q=
+github.com/google/flatbuffers v25.12.19+incompatible h1:haMV2JRRJCe1998HeW/p0X9UaMTK6SDo0ffLn2+DbLs=
+github.com/google/flatbuffers v25.12.19+incompatible/go.mod h1:1AeVuKshWv4vARoZatz6mlQ0JxURH0Kv5+zNeJKJCa8=
+github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8=
+github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU=
+github.com/google/uuid v1.6.0 h1:NIvaJDMOsjHA8n1jAhLSgzrAzy1Hgr+hNrb57e+94F0=
+github.com/google/uuid v1.6.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo=
+github.com/klauspost/asmfmt v1.3.2 h1:4Ri7ox3EwapiOjCki+hw14RyKk201CN4rzyCJRFLpK4=
+github.com/klauspost/asmfmt v1.3.2/go.mod h1:AG8TuvYojzulgDAMCnYn50l/5QV3Bs/tp6j0HLHbNSE=
+github.com/klauspost/compress v1.18.2 h1:iiPHWW0YrcFgpBYhsA6D1+fqHssJscY/Tm/y2Uqnapk=
+github.com/klauspost/compress v1.18.2/go.mod h1:R0h/fSBs8DE4ENlcrlib3PsXS61voFxhIs2DeRhCvJ4=
+github.com/klauspost/cpuid/v2 v2.3.0 h1:S4CRMLnYUhGeDFDqkGriYKdfoFlDnMtqTiI/sFzhA9Y=
+github.com/klauspost/cpuid/v2 v2.3.0/go.mod h1:hqwkgyIinND0mEev00jJYCxPNVRVXFQeu1XKlok6oO0=
+github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8 h1:AMFGa4R4MiIpspGNG7Z948v4n35fFGB3RR3G/ry4FWs=
+github.com/minio/asm2plan9s v0.0.0-20200509001527-cdd76441f9d8/go.mod h1:mC1jAcsrzbxHt8iiaC+zU4b1ylILSosueou12R++wfY=
+github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3 h1:+n/aFZefKZp7spd8DFdX7uMikMLXX4oubIzJF4kv/wI=
+github.com/minio/c2goasm v0.0.0-20190812172519-36a3d3bbc4f3/go.mod h1:RagcQ7I8IeTMnF8JTXieKnO4Z6JCsikNEzj0DwauVzE=
+github.com/pierrec/lz4/v4 v4.1.23 h1:oJE7T90aYBGtFNrI8+KbETnPymobAhzRrR8Mu8n1yfU=
+github.com/pierrec/lz4/v4 v4.1.23/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4=
+github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U=
+github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4=
+github.com/stretchr/objx v0.5.2 h1:xuMeJ0Sdp5ZMRXx/aWO6RZxdr3beISkG5/G/aIRr3pY=
+github.com/stretchr/objx v0.5.2/go.mod h1:FRsXN1f5AsAjCGJKqEizvkpNtU+EGNCLh3NxZ/8L+MA=
+github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U=
+github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U=
+github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU=
+github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E=
+github.com/zeebo/assert v1.3.0 h1:g7C04CbJuIDKNPFHmsk4hwZDO5O+kntRxzaUoNXj+IQ=
+github.com/zeebo/assert v1.3.0/go.mod h1:Pq9JiuJQpG8JLJdtkwrJESF0Foym2/D9XMU5ciN/wJ0=
+github.com/zeebo/xxh3 v1.0.2 h1:xZmwmqxHZA8AI603jOQ0tMqmBr9lPeFwGg6d+xy9DC0=
+github.com/zeebo/xxh3 v1.0.2/go.mod h1:5NWz9Sef7zIDm2JHfFlcQvNekmcEl9ekUZQQKCYaDcA=
+go.opentelemetry.io/auto/sdk v1.2.1 h1:jXsnJ4Lmnqd11kwkBV2LgLoFMZKizbCi5fNZ/ipaZ64=
+go.opentelemetry.io/auto/sdk v1.2.1/go.mod h1:KRTj+aOaElaLi+wW1kO/DZRXwkF4C5xPbEe3ZiIhN7Y=
+go.opentelemetry.io/otel v1.38.0 h1:RkfdswUDRimDg0m2Az18RKOsnI8UDzppJAtj01/Ymk8=
+go.opentelemetry.io/otel v1.38.0/go.mod h1:zcmtmQ1+YmQM9wrNsTGV/q/uyusom3P8RxwExxkZhjM=
+go.opentelemetry.io/otel/metric v1.38.0 h1:Kl6lzIYGAh5M159u9NgiRkmoMKjvbsKtYRwgfrA6WpA=
+go.opentelemetry.io/otel/metric v1.38.0/go.mod h1:kB5n/QoRM8YwmUahxvI3bO34eVtQf2i4utNVLr9gEmI=
+go.opentelemetry.io/otel/sdk v1.38.0 h1:l48sr5YbNf2hpCUj/FoGhW9yDkl+Ma+LrVl8qaM5b+E=
+go.opentelemetry.io/otel/sdk v1.38.0/go.mod h1:ghmNdGlVemJI3+ZB5iDEuk4bWA3GkTpW+DOoZMYBVVg=
+go.opentelemetry.io/otel/sdk/metric v1.38.0 h1:aSH66iL0aZqo//xXzQLYozmWrXxyFkBJ6qT5wthqPoM=
+go.opentelemetry.io/otel/sdk/metric v1.38.0/go.mod h1:dg9PBnW9XdQ1Hd6ZnRz689CbtrUp0wMMs9iPcgT9EZA=
+go.opentelemetry.io/otel/trace v1.38.0 h1:Fxk5bKrDZJUH+AMyyIXGcFAPah0oRcT+LuNtJrmcNLE=
+go.opentelemetry.io/otel/trace v1.38.0/go.mod h1:j1P9ivuFsTceSWe1oY+EeW3sc+Pp42sO++GHkg4wwhs=
+golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0 h1:R84qjqJb5nVJMxqWYb3np9L5ZsaDtB+a39EqjV0JSUM=
+golang.org/x/exp v0.0.0-20250408133849-7e4ce0ab07d0/go.mod h1:S9Xr4PYopiDyqSyp5NjCrhFrqg6A5zA2E/iPHPhqnS8=
+golang.org/x/mod v0.32.0 h1:9F4d3PHLljb6x//jOyokMv3eX+YDeepZSEo3mFJy93c=
+golang.org/x/mod v0.32.0/go.mod h1:SgipZ/3h2Ci89DlEtEXWUk/HteuRin+HHhN+WbNhguU=
+golang.org/x/net v0.49.0 h1:eeHFmOGUTtaaPSGNmjBKpbng9MulQsJURQUAfUwY++o=
+golang.org/x/net v0.49.0/go.mod h1:/ysNB2EvaqvesRkuLAyjI1ycPZlQHM3q01F02UY/MV8=
+golang.org/x/sync v0.19.0 h1:vV+1eWNmZ5geRlYjzm2adRgW2/mcpevXNg50YZtPCE4=
+golang.org/x/sync v0.19.0/go.mod h1:9KTHXmSnoGruLpwFjVSX0lNNA75CykiMECbovNTZqGI=
+golang.org/x/sys v0.40.0 h1:DBZZqJ2Rkml6QMQsZywtnjnnGvHza6BTfYFWY9kjEWQ=
+golang.org/x/sys v0.40.0/go.mod h1:OgkHotnGiDImocRcuBABYBEXf8A9a87e/uXjp9XT3ks=
+golang.org/x/telemetry v0.0.0-20260109210033-bd525da824e2 h1:O1cMQHRfwNpDfDJerqRoE2oD+AFlyid87D40L/OkkJo=
+golang.org/x/telemetry v0.0.0-20260109210033-bd525da824e2/go.mod h1:b7fPSJ0pKZ3ccUh8gnTONJxhn3c/PS6tyzQvyqw4iA8=
+golang.org/x/text v0.33.0 h1:B3njUFyqtHDUI5jMn1YIr5B0IE2U0qck04r6d4KPAxE=
+golang.org/x/text v0.33.0/go.mod h1:LuMebE6+rBincTi9+xWTY8TztLzKHc/9C1uBCG27+q8=
+golang.org/x/tools v0.41.0 h1:a9b8iMweWG+S0OBnlU36rzLp20z1Rp10w+IY2czHTQc=
+golang.org/x/tools v0.41.0/go.mod h1:XSY6eDqxVNiYgezAVqqCeihT4j1U2CCsqvH3WhQpnlg=
+golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da h1:noIWHXmPHxILtqtCOPIhSt0ABwskkZKjD3bXGnZGpNY=
+golang.org/x/xerrors v0.0.0-20240903120638-7835f813f4da/go.mod h1:NDW/Ps6MPRej6fsCIbMTohpP40sJ/P/vI1MoTEGwX90=
+gonum.org/v1/gonum v0.16.0 h1:5+ul4Swaf3ESvrOnidPp4GZbzf0mxVQpDCYUQE7OJfk=
+gonum.org/v1/gonum v0.16.0/go.mod h1:fef3am4MQ93R2HHpKnLk4/Tbh/s0+wqD5nfa6Pnwy4E=
+google.golang.org/genproto/googleapis/rpc v0.0.0-20251029180050-ab9386a59fda h1:i/Q+bfisr7gq6feoJnS/DlpdwEL4ihp41fvRiM3Ork0=
+google.golang.org/genproto/googleapis/rpc v0.0.0-20251029180050-ab9386a59fda/go.mod h1:7i2o+ce6H/6BluujYR+kqX3GKH+dChPTQU19wjRPiGk=
+google.golang.org/grpc v1.78.0 h1:K1XZG/yGDJnzMdd/uZHAkVqJE+xIDOcmdSFZkBUicNc=
+google.golang.org/grpc v1.78.0/go.mod h1:I47qjTo4OKbMkjA/aOOwxDIiPSBofUtQUI5EfpWvW7U=
+google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE=
+google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco=
+gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=
+gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM=
diff --git a/dataset/reader.go b/dataset/reader.go
new file mode 100644
index 0000000..e882332
--- /dev/null
+++ b/dataset/reader.go
@@ -0,0 +1,135 @@
+package dataset
+
+import (
+ "context"
+ "errors"
+ "fmt"
+ "iter"
+ "path/filepath"
+ "slices"
+
+ "github.com/apache/arrow-go/v18/arrow"
+ "github.com/apache/arrow-go/v18/arrow/array"
+ "github.com/apache/arrow-go/v18/arrow/memory"
+ "github.com/apache/arrow-go/v18/parquet/file"
+ "github.com/apache/arrow-go/v18/parquet/pqarrow"
+)
+
+type Reader struct {
+ shards []string
+ batchSize int64
+ err error
+}
+
+func NewReader(name string) (*Reader, error) {
+ var shards []string
+
+ if matches, err := filepath.Glob(filepath.Join(name, "*.parquet")); err != nil {
+ return nil, err
+ } else {
+ shards = matches
+ }
+
+ if len(shards) == 0 {
+ return nil, fmt.Errorf("no parquet files in %q", name)
+ }
+
+ slices.Sort(shards)
+
+ r := &Reader{
+ shards: shards,
+ batchSize: 1024,
+ }
+
+ return r, nil
+}
+
+func (r *Reader) Err() error {
+ return r.err
+}
+
+func (r *Reader) Texts(column string) iter.Seq[string] {
+ return func(yield func(string) bool) {
+ for _, name := range r.shards {
+ err := read(name, column, r.batchSize, yield)
+
+ if errors.Is(err, stop) {
+ return
+ }
+
+ if err != nil {
+ r.err = err
+
+ return
+ }
+ }
+ }
+}
+
+var stop = errors.New("stop")
+
+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 {
+ return err
+ } else {
+ reader = r
+
+ defer reader.Close()
+ }
+
+ var arrowReader *pqarrow.FileReader
+
+ if r, err := pqarrow.NewFileReader(reader, pqarrow.ArrowReadProperties{BatchSize: batchSize}, memory.DefaultAllocator); err != nil {
+ return err
+ } else {
+ arrowReader = r
+ }
+
+ var schema *arrow.Schema
+
+ if s, err := arrowReader.Schema(); err != nil {
+ return err
+ } else {
+ schema = s
+ }
+
+ idxs := schema.FieldIndices(column)
+
+ if len(idxs) == 0 {
+ return errors.New("unknown column")
+ }
+
+ var recordReader pqarrow.RecordReader
+
+ if r, err := arrowReader.GetRecordReader(context.TODO(), []int{idxs[0]}, nil); err != nil {
+ return err
+ } else {
+ recordReader = r
+
+ defer recordReader.Release()
+ }
+
+ for recordReader.Next() {
+ record := recordReader.RecordBatch()
+
+ text, ok := record.Column(0).(*array.String)
+
+ if !ok {
+ return fmt.Errorf("unexpected column type")
+ }
+
+ for i := range int(record.NumRows()) {
+ if !text.IsValid(i) {
+ continue
+ }
+
+ if !yield(text.Value(i)) {
+ return stop
+ }
+ }
+ }
+
+ return nil
+}