summaryrefslogtreecommitdiff
path: root/llm/perplexity.go
diff options
context:
space:
mode:
Diffstat (limited to 'llm/perplexity.go')
-rw-r--r--llm/perplexity.go247
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
-}