diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-17 20:16:39 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-03-17 20:16:39 +0100 |
| commit | 57fb765bdc5819f8a367ad75c731a50da68b8d85 (patch) | |
| tree | 730647fc08847624c138a83f3e8329b6f4755ea6 /llm/perplexity.go | |
| parent | 035eef3308edd4ec581061440337aaa10a37f573 (diff) | |
Update perplexity evaluator
* Introduce results channel
* Introduce jobs channel
* Generalize evaluation via callback
Diffstat (limited to 'llm/perplexity.go')
| -rw-r--r-- | llm/perplexity.go | 247 |
1 files changed, 87 insertions, 160 deletions
diff --git a/llm/perplexity.go b/llm/perplexity.go index 2240f2d..ab31889 100644 --- a/llm/perplexity.go +++ b/llm/perplexity.go @@ -2,234 +2,161 @@ package llm import ( "context" - "log" - "math" - "slices" - "strings" + "fmt" "sync" "time" - "github.com/jonasknobloch/mbpe" - "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/tui" ) -func (e *Evaluator) Perplexity(data *dataset.ParquetReader, window, stride int) (float64, error) { - docs := make([]string, 0) - - n := 0 - - for d := range data.Texts("text") { - if n > 5 { - break - } - - docs = append(docs, d) - - n++ - } - - tokens := toInt64(e.tokenizer.Tokenize(strings.Join(docs, "\n\n")))[:10240] // TODO performance - - batchSize := 1 +func (e *Evaluator[R]) Run(data dataset.Reader, window, stride int) error { + devices := make([]int, len(e.models)) - if len(tokens) < window { - return 0, nil // TODO handle + for i := range len(devices) { + devices[i] = i } - windows := ((len(tokens) - window) / stride) + 1 - jobs := (windows + batchSize - 1) / batchSize + devicePool := newPool(devices...) - pb := mbpe.NewProgressBar("Perplexity", 20, jobs, time.Now()) - - ctx, cancel := context.WithCancel(context.Background()) - - defer cancel() + var wg sync.WaitGroup - done := make(chan struct{}) + for range e.numWorkers { + wg.Add(1) - go func(ctx context.Context) { - main: - for { - select { - case <-ctx.Done(): - break main - default: - time.Sleep(time.Second * 1) + go func() { + defer wg.Done() - e.mutex.RLock() + for b := range e.jobs { + device := devicePool.Acquire() - j := e.jobs + func() { + defer devicePool.Release(device) - e.mutex.RUnlock() + defer func() { + if r := recover(); r != nil { + fmt.Println("HOUSTON") // TODO handle + } + }() - pb.Update(j) - pb.Print() + e.execute(&b, device) + }() - if j >= jobs { - break main - } + e.completed.Add(int64(b.Size())) } - } - - pb.Finish() - - close(done) - }(ctx) - - if err := e.schedule(tokens, window, stride, batchSize); err != nil { - log.Fatal(err) + }() } - <-done + ctx, cancel := context.WithCancel(context.Background()) - totalNLL := float64(0) - totalTokens := 0 + defer cancel() - for _, r := range e.results { - totalNLL += r.nll - totalTokens += r.n - } + pb := tui.NewProgressBar("Perplexity", 20, 0, time.Now()) - average := totalNLL / float64(totalTokens) + go pb.Watch(ctx, 1*time.Second, func() int { + return int(e.completed.Load()) + }) - return math.Exp(average), nil -} + n := 0 -func (e *Evaluator) schedule(tokens []int64, contextSize, stride, batchSize int) error { - jobs := make(chan *job) + for d := range data.Texts("text") { + tokens := toInt64(e.tokenizer.Tokenize(d)) - var wg sync.WaitGroup + e.schedule(n, tokens, window, stride, e.batchSize) - devices := make([]int, len(e.models)) + pb.SetTotal(pb.Total() + e.estimateJobs(tokens, window, stride)) - for i := range len(devices) { - devices[i] = i + n++ } - devicePool := newPool[int](devices...) + close(e.jobs) - for d := 0; d < devicePool.Len(); d++ { - wg.Add(1) - - go func() { - defer wg.Done() - - for j := range jobs { - device := devicePool.Acquire() - - e.execute(j, device) - - devicePool.Release(device) - - e.mutex.Lock() + wg.Wait() - for _, p := range j.results { - e.results = append(e.results, p) - } + close(e.results) - e.jobs++ + return nil +} - e.mutex.Unlock() - } - }() +func (e *Evaluator[R]) estimateJobs(tokens []int64, window, stride int) int { + if len(tokens) < window { + return 0 } - j := newJob(batchSize) + windows := ((len(tokens) - window) / stride) + 1 + // jobs := (windows + batchSize - 1) / e.batchSize - // 0 to 1023: full logits - // 512 to 1535: 1024 upwards - // 1024 to 2047: 1536 upwards - // ... + 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 < len(tokens); i += stride { - for i := 0; i+contextSize <= len(tokens); i += stride { - // if (len(j.positions) == batchSize) || i+stride > len(tokens) { - if len(j.positions) == batchSize { - jobs <- j + // for i := 0; i+contextSize <= len(tokens); i += stride { + for i := 0; i < len(tokens); i += stride { + if b.Size() == batchSize { + e.jobs <- *b - j = newJob(batchSize) + b = newBatch(batchSize) } - j.positions = append(j.positions, n) - j.tokens = append(j.tokens, tokens[i:min(i+contextSize, len(tokens))]) // TODO verify - j.seen = append(j.seen, seen-i) + 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 len(j.positions) > 0 { - jobs <- j + if b.Size() > 0 { + e.jobs <- *b } - close(jobs) - - wg.Wait() - - return nil + e.scheduled.Add(int64(n)) } -func (e *Evaluator) execute(j *job, device int) { - if len(j.positions) != 1 { +func (e *Evaluator[R]) execute(j *batch, device int) { + if j.Size() != 1 { panic("unimplemented") } - if j.seen[0] < 1 { + job := j.jobs[0] + + if job.Seen < 1 { panic("empty context") } m := e.models[device] - logits := make([][]float32, 0, len(j.tokens[0])) - - // fmt.Println("executing job", j.positions[0]) + logits := make([][]float32, 0, len(job.Tokens)) - if _, err := m.Generate(j.tokens[0], 0, &logits); err != nil { + if _, err := m.Generate(job.Tokens, 0, &logits); err != nil { panic(err) // TODO handle } - nllLogits := logits[j.seen[0]-1 : len(logits)-1] - nnlTargets := toInt(j.tokens[0][j.seen[0]:]) + l := logits[job.Seen-1 : len(logits)-1] + t := toInt(job.Tokens[job.Seen:]) - nll := negLogLikelihood(nllLogits, nnlTargets) + r := e.callback(job, l, t) - j.results = append(j.results, result{ - nll: nll, - n: len(nnlTargets), - }) + e.results <- r return } - -func negLogLikelihood(logits [][]float32, targets []int) float64 { - if len(logits) != len(targets) { - panic("mismatched input lengths") - } - - total := float64(0) - - for i, target := range targets { - maxLogit := float64(slices.Max(logits[i])) - - sumExp := float64(0) - - for _, v := range logits[i] { - sumExp += math.Exp(float64(v) - maxLogit) - } - - logSumExp := maxLogit + math.Log(sumExp) - - targetLogit := float64(logits[i][target]) - - logProb := targetLogit - logSumExp - - total -= logProb - } - - return total -} |
