summaryrefslogtreecommitdiff
path: root/gpt2/cmd
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2026-04-04 22:14:02 +0200
committerJonas Knobloch <jonas.knobloch@t-online.de>2026-04-04 22:14:02 +0200
commit60b71d9f78895d78a641b1109e2b0ca293835698 (patch)
tree0057c3c35362e7b6489476be02e8e6a0edfa9619 /gpt2/cmd
parenta96c1bf68bf5c477e1d3f105c18346c80d80b47a (diff)
Add score method
Diffstat (limited to 'gpt2/cmd')
-rw-r--r--gpt2/cmd/gpt2/main.go77
1 files changed, 74 insertions, 3 deletions
diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go
index cad3333..6c7a46b 100644
--- a/gpt2/cmd/gpt2/main.go
+++ b/gpt2/cmd/gpt2/main.go
@@ -3,26 +3,97 @@ package main
import (
"fmt"
"log"
+ "math"
+ "slices"
"go.jknobloc.com/x/gpt2"
)
func main() {
- prompt := []int64{464, 2068, 7586, 21831}
+ prompt := []int64{464, 2068, 7586}
- m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig())
+ generate(prompt) // [-13.483142 -11.277906]
+ score(prompt) // [-13.48314 -11.277912]
+
+ _ = prompt
+}
+
+func generate(prompt []int64) {
+ m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_cache.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), true, false)
if err := m.Init(); err != nil {
log.Fatal(err)
}
- if out, err := m.Generate(prompt, 5, nil); err != nil {
+ logits := make([][]float32, 0)
+
+ if out, err := m.Generate(prompt, 0, &logits); err != nil {
log.Fatal(err)
} else {
fmt.Printf("\n%v\n", out)
}
+ fmt.Println(selectLogProbs(logits[:len(logits)-1], prompt[1:]))
+
+ if err := m.Destroy(); err != nil {
+ log.Fatal(err)
+ }
+}
+
+func score(prompt []int64) {
+ m := gpt2.NewModel("gpt2/models/mbpe_conv/gpt2_8192_m000_babylm_v2/model_eval.onnx", "0", gpt2.NewDefaultConfig().WithVocabSize(8193), false, true)
+
+ if err := m.Init(); err != nil {
+ log.Fatal(err)
+ }
+
+ logProbs := make([]float32, 0, 2)
+
+ if err := m.Score(prompt, &logProbs); err != nil {
+ log.Fatal(err)
+ }
+
+ fmt.Println(logProbs)
+
if err := m.Destroy(); err != nil {
log.Fatal(err)
}
}
+
+func selectLogProbs(logits [][]float32, tokens []int64) []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
+}