summaryrefslogtreecommitdiff
path: root/gpt2/allocator.go
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/allocator.go
parenta96c1bf68bf5c477e1d3f105c18346c80d80b47a (diff)
Add score method
Diffstat (limited to 'gpt2/allocator.go')
-rw-r--r--gpt2/allocator.go74
1 files changed, 63 insertions, 11 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")