summaryrefslogtreecommitdiff
path: root/llmc/tokenize.go
diff options
context:
space:
mode:
Diffstat (limited to 'llmc/tokenize.go')
-rw-r--r--llmc/tokenize.go43
1 files changed, 14 insertions, 29 deletions
diff --git a/llmc/tokenize.go b/llmc/tokenize.go
index 8237af9..742c4be 100644
--- a/llmc/tokenize.go
+++ b/llmc/tokenize.go
@@ -5,48 +5,33 @@ import (
"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
-
+ texts := func(yield func(string) bool) {
for _, text := range reader.Texts() {
- sem <- struct{}{}
-
- wg.Add(1)
-
- go func(t string) {
- defer func() { <-sem; wg.Done() }()
-
- tokens := tok.Tokenize(t)
+ if !yield(text) {
+ return
+ }
+ }
+ }
- tokensU32 := make([]uint32, len(tokens)+1)
+ return Pool(texts, max(1, runtime.NumCPU()-1), func(text string) []uint32 {
+ ids := tok.Tokenize(text)
- tokensU32[0] = uint32(eot)
+ tokens := make([]uint32, 0, len(ids)+1)
- for i, id := range tokens {
- tokensU32[i+1] = uint32(id)
- }
+ tokens = append(tokens, uint32(eot))
- out <- tokensU32
- }(text)
+ for _, id := range ids {
+ tokens = append(tokens, uint32(id))
}
- wg.Wait()
- }()
-
- return out
+ return tokens
+ })
}
func WriteShards(name, data string, shardSize int, docs <-chan []uint32) error {