diff options
Diffstat (limited to 'dataset/count_test.go')
| -rw-r--r-- | dataset/count_test.go | 135 |
1 files changed, 135 insertions, 0 deletions
diff --git a/dataset/count_test.go b/dataset/count_test.go new file mode 100644 index 0000000..0df1ffe --- /dev/null +++ b/dataset/count_test.go @@ -0,0 +1,135 @@ +package dataset + +import ( + "bufio" + "os" + "sync" + "testing" + + "go.jknobloc.com/x/shelf" +) + +func TestCountLinesAll(t *testing.T) { + names := []string{ + shelf.Abs("data/babylm/train_100M/bnc_spoken.train"), + shelf.Abs("data/babylm/train_100M/childes.train"), + shelf.Abs("data/babylm/train_100M/gutenberg.train"), + shelf.Abs("data/babylm/train_100M/open_subtitles.train"), + shelf.Abs("data/babylm/train_100M/simple_wiki.train"), + shelf.Abs("data/babylm/train_100M/switchboard.train"), + } + + n, err := countLinesAll(names, []byte("\n")) + + if err != nil { + t.Fatal(err) + } + + m, err := countLinesNaive(names...) + + if err != nil { + t.Fatal(err) + } + + if n != m { + t.Errorf("expected %d but got %d\n", m, n) + } +} + +func BenchmarkCountAll(b *testing.B) { + names := []string{ + shelf.Abs("data/babylm/train_100M/bnc_spoken.train"), + shelf.Abs("data/babylm/train_100M/childes.train"), + shelf.Abs("data/babylm/train_100M/gutenberg.train"), + shelf.Abs("data/babylm/train_100M/open_subtitles.train"), + shelf.Abs("data/babylm/train_100M/simple_wiki.train"), + shelf.Abs("data/babylm/train_100M/switchboard.train"), + } + + for i := 0; i < b.N; i++ { + _, err := countLinesAll(names, []byte("\n")) + + if err != nil { + b.Fatal(err) + } + } +} + +func BenchmarkCountLinesNaive(b *testing.B) { + names := []string{ + shelf.Abs("data/babylm/train_100M/bnc_spoken.train"), + shelf.Abs("data/babylm/train_100M/childes.train"), + shelf.Abs("data/babylm/train_100M/gutenberg.train"), + shelf.Abs("data/babylm/train_100M/open_subtitles.train"), + shelf.Abs("data/babylm/train_100M/simple_wiki.train"), + shelf.Abs("data/babylm/train_100M/switchboard.train"), + } + + for i := 0; i < b.N; i++ { + _, err := countLinesNaive(names...) + + if err != nil { + b.Fatal(err) + } + } +} + +func countLinesNaive(names ...string) (int, error) { + var wg sync.WaitGroup + + results := make(chan int, len(names)) + errors := make(chan error, len(names)) + + for _, name := range names { + wg.Add(1) + + go func() { + defer wg.Done() + + var scanner *bufio.Scanner + + if file, err := os.Open(name); err != nil { + results <- 0 + errors <- err + + return + } else { + scanner = bufio.NewScanner(file) + + buf := make([]byte, 0, 1024*1024) + + scanner.Buffer(buf, 1024*1024) + + defer file.Close() + } + + count := 0 + + for scanner.Scan() { + count++ + } + + if err := scanner.Err(); err != nil { + results <- 0 + errors <- err + + return + } + + results <- count + }() + } + + wg.Wait() + + close(results) + close(errors) + + total := 0 + + for count := range results { + total += count + } + + return total, <-errors +} |
