diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-25 21:50:15 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-25 21:50:15 +0100 |
| commit | a52a3cd0e6d6d999e93877150adf4c4ec19deff0 (patch) | |
| tree | f10b6b1a35cb504d6b467b9e8adc87376282bbfd /llm/run.go | |
| parent | 0c6a2cadf261945558c75256d32fadf569b4a48d (diff) | |
Refactor llm module
Diffstat (limited to 'llm/run.go')
| -rw-r--r-- | llm/run.go | 159 |
1 files changed, 159 insertions, 0 deletions
diff --git a/llm/run.go b/llm/run.go new file mode 100644 index 0000000..df48ff6 --- /dev/null +++ b/llm/run.go @@ -0,0 +1,159 @@ +package llm + +import ( + "fmt" + "sync" + "time" + + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/tui" +) + +func (e *Evaluator[R]) Run(title string, data dataset.Reader, window, stride int) error { + devices := make([]int, len(e.models)) + + for i := range len(devices) { + devices[i] = i + } + + devicePool := newPool(devices...) + + var wg sync.WaitGroup + + for range e.numWorkers { + wg.Add(1) + + go func() { + defer wg.Done() + + for b := range e.jobs { + device := devicePool.Acquire() + + func() { + defer devicePool.Release(device) + + defer func() { + if r := recover(); r != nil { + fmt.Println("HOUSTON") // TODO handle + } + }() + + e.execute(&b, device) + }() + + e.completed.Add(int64(b.Size())) + } + }() + } + + pb := tui.NewProgressBar(title, 20, 0, time.Now()) + + pb.Start(1*time.Second, func() int { + return int(e.completed.Load()) + }) + + defer pb.Close() + + n := 0 + + for d := range data.Texts("text") { + tokens := toInt64(e.tokenizer.Tokenize(d)) + + e.schedule(n, tokens, window, stride, e.batchSize) + + pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride)) + + n++ + } + + close(e.jobs) + + wg.Wait() + + close(e.results) + + return nil +} + +func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int { + if len(tokens) < window { + return 0 + } + + windows := ((len(tokens) - window) / stride) + 1 + // jobs := (windows + batchSize - 1) / e.batchSize + + return windows +} + +func (e *Evaluator[R]) schedule(uid int, tokens []int64, contextSize, stride, batchSize int) { + b := newBatch(batchSize) + + seen := 1 // first token as context + n := 0 + + // for i := 0; i+contextSize <= len(tokens); i += stride { + for i := 0; i < len(tokens); i += stride { + if b.Size() == batchSize { + e.jobs <- *b + + b = newBatch(batchSize) + } + + j := min(i+contextSize, len(tokens)) + + if j-i < contextSize { + break // don't add jobs with partial windows + } + + // if j < seen { + // break // don't add jobs with no new tokens + // } + + b.AddJob(Job{ + Document: uid, + Position: n, + Tokens: tokens[i:j], + Seen: seen - i, + }) + + seen = i + contextSize + + n++ + } + + if b.Size() > 0 { + e.jobs <- *b + } + + e.scheduled.Add(int64(n)) +} + +func (e *Evaluator[R]) execute(j *batch, device int) { + if j.Size() != 1 { + panic("unimplemented") + } + + job := j.jobs[0] + + if job.Seen < 1 { + panic("empty context") + } + + m := e.models[device] + + logits := make([][]float32, 0, len(job.Tokens)) + + if _, err := m.Generate(job.Tokens, 0, &logits); err != nil { + panic(err) // TODO handle + } + + l := logits[job.Seen-1 : len(logits)-1] + t := toInt(job.Tokens[job.Seen:]) + + r := e.callback(job, l, t) + + e.results <- r + + return +} |
