summaryrefslogtreecommitdiff
path: root/llm/cmd/eval/logprobs.go
diff options
context:
space:
mode:
Diffstat (limited to 'llm/cmd/eval/logprobs.go')
-rw-r--r--llm/cmd/eval/logprobs.go46
1 files changed, 2 insertions, 44 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
-}