blob: 228e9f2b5783772448165f17335d66556cc4bda1 (
plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
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)
}
|