summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-30 13:10:03 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-30 13:19:08 +0200
commit806dcb7f2c01514275dab76116d494345ffb0fd0 (patch)
treed436e8ca6fd71942d72d37d403f78b7abdf52a1b
parent4740e5aa19012e5f8e67ac8cb857d9741208c364 (diff)
Return number of written shards
-rw-r--r--llmc/cmd/data/fineweb.go2
-rw-r--r--llmc/tokenize.go24
2 files changed, 14 insertions, 12 deletions
diff --git a/llmc/cmd/data/fineweb.go b/llmc/cmd/data/fineweb.go
index a1bf7df..599089d 100644
--- a/llmc/cmd/data/fineweb.go
+++ b/llmc/cmd/data/fineweb.go
@@ -37,7 +37,7 @@ func fineWeb() {
docs := llmc.TokenizeAll(reader, tokenizer, 50256)
- if err := llmc.WriteShards(shelf.Abs("llmc/edu_fineweb100B"), "edu_fineweb", 100_000_000, docs); err != nil {
+ if _, err := llmc.WriteShards(shelf.Abs("llmc/edu_fineweb100B"), "edu_fineweb", 100_000_000, docs); err != nil {
log.Fatal(err)
}
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
}