summaryrefslogtreecommitdiff
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
parenta96c1bf68bf5c477e1d3f105c18346c80d80b47a (diff)
Add score method
-rw-r--r--gpt2/allocator.go74
-rw-r--r--gpt2/cmd/gpt2/main.go77
-rw-r--r--gpt2/model.go52
-rw-r--r--gpt2/model_test.go2
-rw-r--r--llm/causal.go1
-rw-r--r--llm/cmd/eval/main.go2
6 files changed, 182 insertions, 26 deletions
diff --git a/gpt2/allocator.go b/gpt2/allocator.go
index d0ec5ec..818fafc 100644
--- a/gpt2/allocator.go
+++ b/gpt2/allocator.go
@@ -7,19 +7,21 @@ import (
)
type Allocator struct {
- config Config
- step int64
- inputNames []string
- outputNames []string
- values map[string]ort.Value
- withCache bool
+ config Config
+ step int64
+ inputNames []string
+ outputNames []string
+ values map[string]ort.Value
+ withCache bool
+ withLogProbs bool
}
-func NewAllocator(config Config, withCache bool) *Allocator {
+func NewAllocator(config Config, withCache bool, withLogProbs bool) *Allocator {
return &Allocator{
- config: config,
- values: make(map[string]ort.Value),
- withCache: withCache,
+ config: config,
+ values: make(map[string]ort.Value),
+ withCache: withCache,
+ withLogProbs: withLogProbs,
}
}
@@ -38,10 +40,20 @@ func (a *Allocator) InputNames() []string {
}
func (a *Allocator) OutputNames() []string {
- names := make([]string, 0, 1+2*a.config.nLayers)
+ capacity := 1 + 2*a.config.nLayers
+
+ if a.withLogProbs {
+ capacity++
+ }
+
+ names := make([]string, 0, capacity)
names = append(names, "logits")
+ if a.withLogProbs {
+ names = append(names, "log_probs")
+ }
+
if a.withCache {
for i := range a.config.nLayers {
names = append(names, fmt.Sprintf("present.%d.key", i), fmt.Sprintf("present.%d.value", i))
@@ -129,6 +141,10 @@ func (a *Allocator) initInputs(tokens []int64) error {
func (a *Allocator) initOutputs(tokens []int64) error {
capacity := 1
+ if a.withLogProbs {
+ capacity++
+ }
+
if a.withCache {
capacity += 2 * a.config.nLayers
}
@@ -141,6 +157,14 @@ func (a *Allocator) initOutputs(tokens []int64) error {
names = append(names, "logits")
+ if a.withLogProbs {
+ if err := a.logProbs(tokens, false); err != nil {
+ return err
+ }
+
+ names = append(names, "log_probs")
+ }
+
if !a.withCache {
a.outputNames = names
@@ -185,6 +209,12 @@ func (a *Allocator) Step(token int64) error {
return err
}
+ if a.withLogProbs {
+ if err := a.logProbs(tokens, true); err != nil {
+ return err
+ }
+ }
+
for i := range int64(a.config.nLayers) {
for _, suffix := range []string{"key", "value"} {
if err := a.rotateCache(tokens, i, suffix); err != nil {
@@ -352,6 +382,28 @@ func (a *Allocator) logits(tokens []int64, force bool) error {
return nil
}
+func (a *Allocator) logProbs(tokens []int64, force bool) error {
+ const name = "log_probs"
+
+ if _, ok := a.values[name]; ok {
+ if !force {
+ panic("log_probs already allocated")
+ }
+
+ _ = a.values[name].Destroy()
+ }
+
+ shape := []int64{1, int64(len(tokens)) - 1}
+
+ if t, err := ort.NewEmptyTensor[float32](shape); err != nil {
+ return err
+ } else {
+ a.values[name] = ort.Value(t)
+ }
+
+ return nil
+}
+
func (a *Allocator) presentKeyValues(tokens []int64, start, i int64, suffix string, force bool) error {
if int(i) > a.config.nLayers {
panic("invalid layer index")
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
+}
diff --git a/gpt2/model.go b/gpt2/model.go
index 6557a61..dd18c80 100644
--- a/gpt2/model.go
+++ b/gpt2/model.go
@@ -12,18 +12,22 @@ import (
)
type Model struct {
- name string
- deviceID string
- config Config
- session *ort.DynamicAdvancedSession
- allocator *Allocator
+ name string
+ deviceID string
+ config Config
+ withCache bool
+ withLogProbs bool
+ session *ort.DynamicAdvancedSession
+ allocator *Allocator
}
-func NewModel(name string, deviceID string, config Config) *Model {
+func NewModel(name string, deviceID string, config Config, withCache bool, withLogProbs bool) *Model {
return &Model{
- name: name,
- deviceID: deviceID,
- config: config,
+ name: name,
+ deviceID: deviceID,
+ config: config,
+ withCache: withCache,
+ withLogProbs: withLogProbs,
}
}
@@ -60,7 +64,7 @@ func (m *Model) Init() error {
return err
}
- m.allocator = NewAllocator(m.config, true)
+ m.allocator = NewAllocator(m.config, m.withCache, m.withLogProbs)
var options *ort.SessionOptions
@@ -100,6 +104,10 @@ func (m *Model) Destroy() error {
}
func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) {
+ if m.withLogProbs {
+ panic("generate called on eval model")
+ }
+
if len(prompt) == 0 {
return nil, errors.New("empty prompt")
}
@@ -151,6 +159,30 @@ func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]in
return r, nil
}
+func (m *Model) Score(tokens []int64, logProbs *[]float32) error {
+ if !m.withLogProbs {
+ panic("score called on default model")
+ }
+
+ if err := m.allocator.Init(tokens); err != nil {
+ return err
+ }
+
+ if err := m.forward(m.allocator); err != nil {
+ return err
+ }
+
+ if logProbs != nil {
+ _, outputs := m.allocator.Outputs()
+
+ d := outputs[1].(*ort.Tensor[float32]).GetData()
+
+ *logProbs = append(*logProbs, d...)
+ }
+
+ return nil
+}
+
func (m *Model) logits(output ort.Value) [][]float32 {
d := output.(*ort.Tensor[float32]).GetData()
n := len(d) / m.config.vocabSize
diff --git a/gpt2/model_test.go b/gpt2/model_test.go
index 18abdb8..6677f73 100644
--- a/gpt2/model_test.go
+++ b/gpt2/model_test.go
@@ -26,7 +26,7 @@ func fromModel() []float32 {
}
func model() *Model {
- m := NewModel("models/base/model.onnx", "0", NewDefaultConfig()) // TODO check if CUDA is available
+ m := NewModel("models/base/model.onnx", "0", NewDefaultConfig(), true, false) // TODO check if CUDA is available
if err := m.Init(); err != nil {
log.Fatal(err)
diff --git a/llm/causal.go b/llm/causal.go
index 296a744..b683985 100644
--- a/llm/causal.go
+++ b/llm/causal.go
@@ -2,4 +2,5 @@ package llm
type Causal interface {
Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error)
+ Score(tokens []int64, logProbs *[]float32) error
}
diff --git a/llm/cmd/eval/main.go b/llm/cmd/eval/main.go
index 50448be..44dfedd 100644
--- a/llm/cmd/eval/main.go
+++ b/llm/cmd/eval/main.go
@@ -26,7 +26,7 @@ func data() *dataset.ParquetReader {
}
func model() *gpt2.Model {
- m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig())
+ m := gpt2.NewModel("gpt2/models/base/model.onnx", "0", gpt2.NewDefaultConfig(), true, false)
if err := m.Init(); err != nil {
log.Fatal(err)