summaryrefslogtreecommitdiff
path: root/llmc/tokenize.go
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-22 15:36:58 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-23 20:17:53 +0200
commitbabd70eb9d4e4b028c0eb8c32e6bed65dafa0481 (patch)
tree75635ad8e0ef6541b7b5d6926f529830b32f715b /llmc/tokenize.go
parent54b11da0d2f853395f270379044674363b1a124a (diff)
Add llmc module
Diffstat (limited to 'llmc/tokenize.go')
-rw-r--r--llmc/tokenize.go113
1 files changed, 113 insertions, 0 deletions
diff --git a/llmc/tokenize.go b/llmc/tokenize.go
new file mode 100644
index 0000000..8237af9
--- /dev/null
+++ b/llmc/tokenize.go
@@ -0,0 +1,113 @@
+package llmc
+
+import (
+ "fmt"
+ "os"
+ "path/filepath"
+ "runtime"
+ "sync"
+
+ "go.jknobloc.com/x/dataset"
+ "go.jknobloc.com/x/llm"
+)
+
+func TokenizeAll(reader dataset.Reader, tok llm.Tokenizer, eot int) <-chan []uint32 {
+ out := make(chan []uint32, 256)
+
+ go func() {
+ defer close(out)
+
+ sem := make(chan struct{}, max(1, runtime.NumCPU()-1))
+
+ var wg sync.WaitGroup
+
+ for _, text := range reader.Texts() {
+ sem <- struct{}{}
+
+ wg.Add(1)
+
+ go func(t string) {
+ defer func() { <-sem; wg.Done() }()
+
+ tokens := tok.Tokenize(t)
+
+ tokensU32 := make([]uint32, len(tokens)+1)
+
+ tokensU32[0] = uint32(eot)
+
+ for i, id := range tokens {
+ tokensU32[i+1] = uint32(id)
+ }
+
+ out <- tokensU32
+ }(text)
+ }
+
+ wg.Wait()
+ }()
+
+ return out
+}
+
+func WriteShards(name, data string, shardSize int, docs <-chan []uint32) error {
+ if err := os.MkdirAll(name, os.ModePerm); err != nil {
+ return err
+ }
+
+ buffer := make([]uint32, 0, shardSize)
+
+ shardIdx := 0
+
+ flush := func() error {
+ split := "train"
+
+ if shardIdx == 0 {
+ split = "val"
+ }
+
+ path := filepath.Join(name, fmt.Sprintf("%s_%s_%06d.bin", data, split, shardIdx))
+
+ d := DataFile[uint32]{
+ Model: GPT2,
+ Tokens: buffer,
+ }
+
+ if n, err := Serialize(&d, path); err != nil {
+ return err
+ } else {
+ fmt.Printf("wrote %s (%d tokens)\n", filepath.Base(path), n)
+ }
+
+ buffer = buffer[:0]
+
+ shardIdx++
+
+ return nil
+ }
+
+ for tokens := range docs {
+ for len(tokens) > 0 {
+ space := shardSize - len(buffer)
+
+ if space >= len(tokens) {
+ buffer = append(buffer, tokens...)
+
+ break
+ }
+
+ buffer = append(buffer, tokens[:space]...)
+
+ tokens = tokens[space:] // carry remainder
+
+ if err := flush(); err != nil {
+ return err
+ }
+ }
+ }
+
+ if len(buffer) > 0 {
+ return flush()
+ }
+
+ return nil
+}