summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorJonas Knobloch <jonas.knobloch@t-online.de>2025-11-11 01:26:57 +0100
committerJonas Knobloch <jonas.knobloch@t-online.de>2025-11-11 01:26:57 +0100
commiteb068eca5af159e0268ec6e06f93c57f9ec08026 (patch)
treeb6d7cd1131cc45ac41f437c1154918bb702e7d87
parent4d14f6119407a3be6d74c80e426646798910dabe (diff)
Support multi-token prompts
-rw-r--r--main.go35
1 files changed, 32 insertions, 3 deletions
diff --git a/main.go b/main.go
index 2cbb28e..2c2e567 100644
--- a/main.go
+++ b/main.go
@@ -1,6 +1,7 @@
package main
import (
+ "errors"
"fmt"
"log"
"math"
@@ -25,16 +26,44 @@ func main() {
defer ort.DestroyEnvironment()
- if err := generate("scripts/onnx-gpt2/model.onnx", 464, 5); err != nil {
+ prompt := []int64{464, 257, 5025}
+
+ if err := generate("scripts/onnx-gpt2/model.onnx", prompt, 5); err != nil {
log.Fatal(err)
}
}
-func generate(model string, token int64, steps int64) error {
+func generate(model string, prompt []int64, steps int64) error {
+ if len(prompt) == 0 {
+ return errors.New("empty prompt")
+ }
+
cacheNames, cacheValues := emptyCache()
+ if len(prompt) > 1 {
+ for i := range prompt {
+ if i == len(prompt)-1 {
+ break
+ }
+
+ position := int64(i)
+ token := prompt[i]
+
+ _, _, outputs, err := forward(model, token, position, cacheNames, cacheValues)
+
+ if err != nil {
+ return err
+ }
+
+ cacheValues = outputs[1:]
+ }
+ }
+
+ token := prompt[len(prompt)-1]
+ offset := int64(len(prompt)) - 1
+
for step := range steps {
- logits, _, outputs, err := forward(model, token, step, cacheNames, cacheValues)
+ logits, _, outputs, err := forward(model, token, offset+step, cacheNames, cacheValues)
if err != nil {
return err