From 9349bdb0c5168caaa20e9a3df932a0c926592374 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Sat, 4 Apr 2026 22:33:22 +0200 Subject: Bind either logits or log probs --- gpt2/allocator.go | 34 +++++++++++++--------------------- 1 file changed, 13 insertions(+), 21 deletions(-) (limited to 'gpt2/allocator.go') diff --git a/gpt2/allocator.go b/gpt2/allocator.go index 818fafc..575fac5 100644 --- a/gpt2/allocator.go +++ b/gpt2/allocator.go @@ -40,18 +40,14 @@ func (a *Allocator) InputNames() []string { } func (a *Allocator) OutputNames() []string { - capacity := 1 + 2*a.config.nLayers - - if a.withLogProbs { - capacity++ - } + capacity := 2*a.config.nLayers + 1 names := make([]string, 0, capacity) - names = append(names, "logits") - if a.withLogProbs { names = append(names, "log_probs") + } else { + names = append(names, "logits") } if a.withCache { @@ -141,28 +137,24 @@ 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 } names := make([]string, 0, capacity) - if err := a.logits(tokens, false); err != nil { - return err - } - - names = append(names, "logits") - if a.withLogProbs { if err := a.logProbs(tokens, false); err != nil { return err } names = append(names, "log_probs") + } else { + if err := a.logits(tokens, false); err != nil { + return err + } + + names = append(names, "logits") } if !a.withCache { @@ -205,14 +197,14 @@ func (a *Allocator) Step(token int64) error { return err } - if err := a.logits(tokens, true); err != nil { - return err - } - if a.withLogProbs { if err := a.logProbs(tokens, true); err != nil { return err } + } else { + if err := a.logits(tokens, true); err != nil { + return err + } } for i := range int64(a.config.nLayers) { -- cgit v1.3.1