summaryrefslogtreecommitdiff
path: root/llmc/tokenize.go
diff options
context:
space:
mode:
Diffstat (limited to 'llmc/tokenize.go')
-rw-r--r--llmc/tokenize.go24
1 files changed, 13 insertions, 11 deletions
diff --git a/llmc/tokenize.go b/llmc/tokenize.go
index 1863157..e90de58 100644
--- a/llmc/tokenize.go
+++ b/llmc/tokenize.go
@@ -34,38 +34,38 @@ func TokenizeAll(reader dataset.Reader, tok llm.Tokenizer, eot int) <-chan []uin
})
}
-func WriteShards(path, name string, size int, docs <-chan []uint32) error {
+func WriteShards(path, name string, size int, docs <-chan []uint32) (int, error) {
if err := os.MkdirAll(path, os.ModePerm); err != nil {
- return err
+ return 0, err
}
buffer := make([]uint32, 0, size)
- shardIdx := 0
+ n := 0
flush := func() error {
split := "train"
- if shardIdx == 0 {
+ if n == 0 {
split = "val"
}
- shardName := filepath.Join(path, fmt.Sprintf("%s_%s_%06d.bin", name, split, shardIdx))
+ shardName := filepath.Join(path, fmt.Sprintf("%s_%s_%06d.bin", name, split, n))
d := DataFile[uint32]{
Model: GPT2,
Tokens: buffer,
}
- if n, err := Serialize(&d, shardName); err != nil {
+ if numTokens, err := Serialize(&d, shardName); err != nil {
return err
} else {
- fmt.Printf("wrote %s (%d tokens)\n", filepath.Base(shardName), n)
+ fmt.Printf("wrote %s (%d tokens)\n", filepath.Base(shardName), numTokens)
}
buffer = buffer[:0]
- shardIdx++
+ n++
return nil
}
@@ -85,14 +85,16 @@ func WriteShards(path, name string, size int, docs <-chan []uint32) error {
tokens = tokens[space:] // carry remainder
if err := flush(); err != nil {
- return err
+ return n, err
}
}
}
if len(buffer) > 0 {
- return flush()
+ if err := flush(); err != nil {
+ return n, err
+ }
}
- return nil
+ return n, nil
}