diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-04 22:14:02 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-04-04 22:14:02 +0200 |
| commit | 60b71d9f78895d78a641b1109e2b0ca293835698 (patch) | |
| tree | 0057c3c35362e7b6489476be02e8e6a0edfa9619 | |
| parent | a96c1bf68bf5c477e1d3f105c18346c80d80b47a (diff) | |
Add score method
| -rw-r--r-- | gpt2/allocator.go | 74 | ||||
| -rw-r--r-- | gpt2/cmd/gpt2/main.go | 77 | ||||
| -rw-r--r-- | gpt2/model.go | 52 | ||||
| -rw-r--r-- | gpt2/model_test.go | 2 | ||||
| -rw-r--r-- | llm/causal.go | 1 | ||||
| -rw-r--r-- | llm/cmd/eval/main.go | 2 |
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) |
