From 60b71d9f78895d78a641b1109e2b0ca293835698 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Sat, 4 Apr 2026 22:14:02 +0200 Subject: Add score method --- gpt2/cmd/gpt2/main.go | 77 +++++++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 74 insertions(+), 3 deletions(-) (limited to 'gpt2/cmd') 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 +} -- cgit v1.3.1