diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-22 12:44:08 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-23 20:18:00 +0200 |
| commit | d3276058dd057178e1218a6b76f2ec42d713d6e8 (patch) | |
| tree | aaa5f499b3c9a76c3fc81cd4d3b50a46069eaee8 | |
| parent | babd70eb9d4e4b028c0eb8c32e6bed65dafa0481 (diff) | |
Fix output order
| -rw-r--r-- | llmc/pool.go | 81 | ||||
| -rw-r--r-- | llmc/tokenize.go | 43 |
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 { |
