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 | |
| parent | 035eef3308edd4ec581061440337aaa10a37f573 (diff) | |
Update perplexity evaluator
* Introduce results channel
* Introduce jobs channel
* Generalize evaluation via callback
Diffstat (limited to 'llm')
| -rw-r--r-- | llm/cmd/eval/main.go (renamed from llm/cmd/ppl/ppl.go) | 19 | ||||
| -rw-r--r-- | llm/cmd/eval/ppl.go | 80 | ||||
| -rw-r--r-- | llm/evaluator.go | 43 | ||||
| -rw-r--r-- | llm/job.go | 37 | ||||
| -rw-r--r-- | llm/perplexity.go | 247 | ||||
| -rw-r--r-- | llm/perplexity_test.go | 45 | ||||
| -rw-r--r-- | llm/pool.go | 1 | ||||
| -rw-r--r-- | llm/result.go | 6 | ||||
| -rw-r--r-- | llm/stat.go | 34 |
9 files changed, 304 insertions, 208 deletions
diff --git a/llm/cmd/ppl/ppl.go b/llm/cmd/eval/main.go index 6912e48..05b3f96 100644 --- a/llm/cmd/ppl/ppl.go +++ b/llm/cmd/eval/main.go @@ -1,33 +1,16 @@ package main import ( - "fmt" "log" "github.com/jonasknobloch/mbpe" "go.jknobloc.com/x/dataset" "go.jknobloc.com/x/gpt2" - "go.jknobloc.com/x/llm" ) func main() { - d := data() - m := model() - t := tokenizer() - - e := llm.NewEvaluator() - - e.SetTokenizer(t) - e.AddModel(m) - - ppl, err := e.Perplexity(d, 1024, 512) - - if err != nil { - log.Fatal(err) - } - - fmt.Println(ppl) + perplexity() } func data() *dataset.ParquetReader { diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go new file mode 100644 index 0000000..fdbcff3 --- /dev/null +++ b/llm/cmd/eval/ppl.go @@ -0,0 +1,80 @@ +package main + +import ( + "fmt" + "log" + "math" + "strings" + "sync" + + "go.jknobloc.com/x/dataset" + "go.jknobloc.com/x/llm" +) + +type pplResult struct { + v float64 + n int +} + +func perplexity() { + d := data() + m := model() + t := tokenizer() + + e := llm.NewEvaluator[pplResult]() + + e.SetTokenizer(t) + e.AddModel(m) + + e.SetCallback(func(job llm.Job, logits [][]float32, tokens []int) pplResult { + p, n := llm.NegLogLikelihood(logits, tokens) + + return pplResult{ + v: p, + n: n, + } + }) + + results := e.Results() + + var wg sync.WaitGroup + + total := float64(0) + n := 0 + + wg.Add(1) + + go func() { + defer wg.Done() + + for r := range results { + total += r.v + n += r.n + } + }() + + if err := e.Run(d, 1024, 512); err != nil { + log.Fatal(err) + } + + wg.Wait() + + avg := total / float64(n) + ppl := math.Exp(avg) + + fmt.Println(ppl) +} + +func joined() dataset.Reader { + miniPile := data() + + docs := make([]string, 0) + + for d := range miniPile.Texts("text") { + docs = append(docs, d) + } + + j := dataset.NewStringReader(strings.Join(docs, "\n\n")) + + return j +} diff --git a/llm/evaluator.go b/llm/evaluator.go index 9faf483..24fe2d4 100644 --- a/llm/evaluator.go +++ b/llm/evaluator.go @@ -1,26 +1,47 @@ package llm -import "sync" +import ( + "sync/atomic" +) -type Evaluator struct { - mutex sync.RWMutex +type Evaluator[R any] struct { models []Causal tokenizer Tokenizer - results []result - jobs int + + batchSize int + numWorkers int + + scheduled atomic.Int64 + completed atomic.Int64 + + jobs chan batch + results chan R + + callback func(job Job, logits [][]float32, tokens []int) R } -func NewEvaluator() *Evaluator { - return &Evaluator{ - models: make([]Causal, 0), - results: make([]result, 0), +func NewEvaluator[R any]() *Evaluator[R] { + return &Evaluator[R]{ + models: make([]Causal, 0), + batchSize: 1, // TODO arg + numWorkers: 4, // TODO arg + jobs: make(chan batch, 1024), + results: make(chan R, 1024), } } -func (e *Evaluator) AddModel(model Causal) { +func (e *Evaluator[R]) AddModel(model Causal) { e.models = append(e.models, model) } -func (e *Evaluator) SetTokenizer(tokenizer Tokenizer) { +func (e *Evaluator[R]) SetTokenizer(tokenizer Tokenizer) { e.tokenizer = tokenizer } + +func (e *Evaluator[R]) SetCallback(callback func(job Job, logits [][]float32, tokens []int) R) { + e.callback = callback +} + +func (e *Evaluator[R]) Results() chan R { + return e.results +} @@ -1,19 +1,30 @@ package llm -type job struct { - positions []int - tokens [][]int64 - seen []int - results []result - debug [][2]int +type Job struct { + Document int + Position int + Tokens []int64 + Seen int } -func newJob(batchSize int) *job { - return &job{ - positions: make([]int, 0, batchSize), - tokens: make([][]int64, 0, batchSize), - seen: make([]int, 0, batchSize), - results: make([]result, 0, batchSize), - debug: make([][2]int, 0), +type batch struct { + jobs []Job +} + +func newBatch(capacity int) *batch { + return &batch{ + jobs: make([]Job, 0, capacity), + } +} + +func (b *batch) Size() int { + return len(b.jobs) +} + +func (b *batch) AddJob(job Job) { + if b.Size() == cap(b.jobs) { + panic("batch full") } + + b.jobs = append(b.jobs, job) } 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 -} diff --git a/llm/perplexity_test.go b/llm/perplexity_test.go new file mode 100644 index 0000000..81f10af --- /dev/null +++ b/llm/perplexity_test.go @@ -0,0 +1,45 @@ +package llm + +import ( + "fmt" + "testing" +) + +func TestEvaluator_estimateJobs(t *testing.T) { + type gold struct { + tokens int + window int + stride int + expected int + } + + tests := []gold{ + {tokens: 20, window: 10, stride: 3, expected: 4}, + {tokens: 20, window: 10, stride: 4, expected: 3}, + {tokens: 20, window: 10, stride: 5, expected: 3}, + + {tokens: 0, window: 1, stride: 1, expected: 0}, + {tokens: 1, window: 1, stride: 1, expected: 1}, + + {tokens: 0, window: 1024, stride: 1, expected: 0}, + {tokens: 1, window: 1, stride: 1024, expected: 1}, + } + + e := NewEvaluator[any]() + + for _, tt := range tests { + t.Run( + fmt.Sprintf("tokens%d_window%d_stride%d", tt.tokens, tt.window, tt.stride), + + func(t *testing.T) { + tokens := make([]int64, tt.tokens) + + got := e.estimateJobs(tokens, tt.window, tt.stride) + + if got != tt.expected { + t.Errorf("expected %d but got (%d, %d, %d) = %d", tt.expected, tt.tokens, tt.window, tt.stride, got) + } + }, + ) + } +} diff --git a/llm/pool.go b/llm/pool.go index 294b64f..68b517d 100644 --- a/llm/pool.go +++ b/llm/pool.go @@ -28,6 +28,7 @@ func (p *pool[K]) Len() int { func (p *pool[K]) Acquire() K { p.mutex.Lock() + defer p.mutex.Unlock() for { diff --git a/llm/result.go b/llm/result.go deleted file mode 100644 index c77f68d..0000000 --- a/llm/result.go +++ /dev/null @@ -1,6 +0,0 @@ -package llm - -type result struct { - nll float64 - n int -} diff --git a/llm/stat.go b/llm/stat.go new file mode 100644 index 0000000..228e9f2 --- /dev/null +++ b/llm/stat.go @@ -0,0 +1,34 @@ +package llm + +import ( + "math" + "slices" +) + +func NegLogLikelihood(logits [][]float32, targets []int) (float64, int) { + 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, len(targets) +} |
