summaryrefslogtreecommitdiff
path: root/llm
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 00:01:23 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-07 00:01:23 +0200
commitb81fa93e098ca689cb6f73627968fe0b4d49decb (patch)
treebb852b92c2761a0924a5cefcc3bce6d10b480822 /llm
parentb9faf92148c7e6432ff0411783f444b7779fa36a (diff)
Run evaluator with score
Diffstat (limited to 'llm')
-rw-r--r--llm/cmd/eval/logprobs.go46
-rw-r--r--llm/cmd/eval/main.go2
-rw-r--r--llm/cmd/eval/ppl.go14
-rw-r--r--llm/evaluator.go4
-rw-r--r--llm/run.go6
-rw-r--r--llm/stat.go34
6 files changed, 19 insertions, 87 deletions
diff --git a/llm/cmd/eval/logprobs.go b/llm/cmd/eval/logprobs.go
index bbc71aa..7228ea6 100644
--- a/llm/cmd/eval/logprobs.go
+++ b/llm/cmd/eval/logprobs.go
@@ -4,8 +4,6 @@ import (
"context"
"database/sql"
"log"
- "math"
- "slices"
_ "github.com/duckdb/duckdb-go/v2"
@@ -27,16 +25,14 @@ func logprobs() {
m := model()
t := tokenizer()
- e := llm.NewEvaluator(m, t, func(j llm.Job, logits [][]float32, tokens []int) logProbs {
- probs := selectLogProbs(logits, tokens)
-
+ e := llm.NewEvaluator(m, t, func(j llm.Job, l []float32, tokens []int) logProbs {
r := make(logProbs, len(tokens))
for i, token := range tokens {
r[i] = logProb{
document: j.Document,
token: token,
- value: probs[i],
+ value: l[i],
offset: i,
}
}
@@ -113,41 +109,3 @@ func insert(stmt *sql.Stmt, prob logProb) error {
return nil
}
-
-func selectLogProbs(logits [][]float32, tokens []int) []float32 {
- if len(logits) != len(tokens) {
- panic("length mismatch")
- }
-
- r := make([]float32, len(tokens))
-
- for i, token := range tokens {
- logprobs := logSoftmax(logits[i])
-
- r[i] = logprobs[token]
- }
-
- return r
-}
-
-func logSoftmax(logits []float32) []float32 {
- m := slices.Max(logits)
-
- s := float32(0.0)
- r := make([]float32, len(logits))
-
- for i, v := range logits {
- e := float32(math.Exp(float64(v - m)))
-
- r[i] = v
- s += e
- }
-
- lse := float32(math.Log(float64(s))) + m
-
- for i := range r {
- r[i] -= lse
- }
-
- return r
-}
diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go
index 43752d9..b654e49 100644
--- a/llm/cmd/eval/main.go
+++ b/llm/cmd/eval/main.go
@@ -34,7 +34,7 @@ func data() *dataset.ParquetReader {
}
func model() *gpt2.Model {
- m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig(), true, true, false)
+ m := gpt2.NewModel("gpt2/models/base/model_eval.onnx", "0", gpt2.NewDefaultConfig(), false, false, true)
if err := m.Init(); err != nil {
log.Fatal(err)
diff --git a/llm/cmd/eval/ppl.go b/llm/cmd/eval/ppl.go
index 261114e..8838f86 100644
--- a/llm/cmd/eval/ppl.go
+++ b/llm/cmd/eval/ppl.go
@@ -21,11 +21,19 @@ func perplexity() {
m := model()
t := tokenizer()
- e := llm.NewEvaluator(m, t, func(job llm.Job, logits [][]float32, tokens []int) pplResult {
- p, n := llm.NegLogLikelihood(logits, tokens)
+ e := llm.NewEvaluator(m, t, func(job llm.Job, logProbs []float32, tokens []int) pplResult {
+ total := float64(0)
+
+ n := 0
+
+ for _, p := range logProbs {
+ total -= float64(p)
+
+ n++
+ }
return pplResult{
- v: p,
+ v: total,
n: n,
}
})
diff --git a/llm/evaluator.go b/llm/evaluator.go
index 17a7759..8340ad4 100644
--- a/llm/evaluator.go
+++ b/llm/evaluator.go
@@ -21,10 +21,10 @@ type Evaluator[R any] struct {
jobs chan batch
results chan R
- callback func(job Job, logits [][]float32, tokens []int) R
+ callback func(job Job, logProbs []float32, tokens []int) R
}
-func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logits [][]float32, tokens []int) R) *Evaluator[R] {
+func NewEvaluator[R any](model Causal, tokenizer Tokenizer, callback func(job Job, logProbs []float32, tokens []int) R) *Evaluator[R] {
return &Evaluator[R]{
models: []Causal{model}, // TODO multiple devices
tokenizer: tokenizer,
diff --git a/llm/run.go b/llm/run.go
index fd77ce9..9cdf8c3 100644
--- a/llm/run.go
+++ b/llm/run.go
@@ -125,13 +125,13 @@ func (e *Evaluator[R]) execute(j *batch, device int) {
m := e.models[device]
- logits := make([][]float32, 0, len(job.Tokens))
+ logProbs := make([]float32, 0, len(job.Tokens)-1)
- if _, err := m.Generate(job.Tokens, 0, &logits); err != nil {
+ if err := m.Score(job.Tokens, &logProbs); err != nil {
panic(err) // TODO handle
}
- l := logits[job.Seen-1 : len(logits)-1]
+ l := logProbs[job.Seen-1:]
t := toInt(job.Tokens[job.Seen:])
r := e.callback(job, l, t)
diff --git a/llm/stat.go b/llm/stat.go
deleted file mode 100644
index 228e9f2..0000000
--- a/llm/stat.go
+++ /dev/null
@@ -1,34 +0,0 @@
-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)
-}