diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2025-11-11 01:26:57 +0100 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2025-11-11 01:26:57 +0100 |
| commit | eb068eca5af159e0268ec6e06f93c57f9ec08026 (patch) | |
| tree | b6d7cd1131cc45ac41f437c1154918bb702e7d87 | |
| parent | 4d14f6119407a3be6d74c80e426646798910dabe (diff) | |
Support multi-token prompts
| -rw-r--r-- | main.go | 35 |
1 files changed, 32 insertions, 3 deletions
@@ -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 |
