summaryrefslogtreecommitdiff
path: root/llmc
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-22 12:44:08 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-23 20:18:00 +0200
commitd3276058dd057178e1218a6b76f2ec42d713d6e8 (patch)
treeaaa5f499b3c9a76c3fc81cd4d3b50a46069eaee8 /llmc
parentbabd70eb9d4e4b028c0eb8c32e6bed65dafa0481 (diff)
Fix output order
Diffstat (limited to 'llmc')
-rw-r--r--llmc/pool.go81
-rw-r--r--llmc/tokenize.go43
2 files changed, 95 insertions, 29 deletions
diff --git a/llmc/pool.go b/llmc/pool.go
new file mode 100644
index 0000000..c73d2d1
--- /dev/null
+++ b/llmc/pool.go
@@ -0,0 +1,81 @@
+package llmc
+
+import (
+ "iter"
+ "sync"
+)
+
+func Pool[In, Out any](seq iter.Seq[In], nWorkers int, f func(In) Out) <-chan Out {
+ out := make(chan Out, 256)
+
+ go func() {
+ defer close(out)
+
+ type work struct {
+ idx int
+ in In
+ }
+
+ type result struct {
+ idx int
+ out Out
+ }
+
+ jobs := make(chan work, nWorkers)
+ results := make(chan result, nWorkers)
+
+ var wg sync.WaitGroup
+
+ for range nWorkers {
+ wg.Add(1)
+
+ go func() {
+ defer wg.Done()
+
+ for j := range jobs {
+ results <- result{j.idx, f(j.in)}
+ }
+ }()
+ }
+
+ go func() {
+ idx := 0
+
+ for item := range seq {
+ jobs <- work{idx, item}
+
+ idx++
+ }
+
+ close(jobs)
+
+ wg.Wait()
+
+ close(results)
+ }()
+
+ next := 0
+
+ pending := make(map[int]Out)
+
+ for r := range results {
+ pending[r.idx] = r.out
+
+ for {
+ v, ok := pending[next]
+
+ if !ok {
+ break
+ }
+
+ out <- v
+
+ delete(pending, next)
+
+ next++
+ }
+ }
+ }()
+
+ return out
+}
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 {