summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
-rw-r--r--main.go51
1 files changed, 18 insertions, 33 deletions
diff --git a/main.go b/main.go
index 82461da..e1f2cdb 100644
--- a/main.go
+++ b/main.go
@@ -26,7 +26,7 @@ func main() {
defer ort.DestroyEnvironment()
- prompt := []int64{464, 257, 5025}
+ prompt := []int64{464, 2068, 7586, 21831}
if out, err := generate("scripts/onnx-gpt2/model.onnx", prompt, 5); err != nil {
log.Fatal(err)
@@ -40,57 +40,42 @@ func generate(model string, prompt []int64, steps int64) ([]int64, error) {
return nil, 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 nil, err
- }
+ context := int64(len(prompt))
- cacheValues = outputs[1:]
- }
- }
+ cacheNames, cacheValues := emptyCache()
- token := prompt[len(prompt)-1]
- offset := int64(len(prompt)) - 1
+ token := prompt[0]
- out := make([]int64, steps)
+ out := make([]int64, 0, steps+1)
- for step := range steps {
- logits, _, outputs, err := forward(model, token, offset+step, cacheNames, cacheValues)
+ for step := range context + steps {
+ _, _, outputs, err := forward(model, token, step, cacheNames, cacheValues)
if err != nil {
return nil, err
}
- probs := softmax(logits.GetData())
+ logits := outputs[0].(*ort.Tensor[float32]).GetData()
- idx, p := topK(probs, 5)
+ idx, p := topK(softmax(logits), 5)
- fmt.Println()
+ fmt.Printf("\n%d\n\n", token)
for i, t := range idx {
- fmt.Printf("%.4f [%d]\n", p[i], t)
+ fmt.Printf("%.4f %.4f [%d]\n", logits[t], p[i], t)
}
- token = int64(idx[0]) // choose best token
-
- out[step] = token
+ if step < context-1 {
+ token = prompt[step+1]
+ } else {
+ token = int64(idx[0]) // choose best token
+ out = append(out, token)
+ }
cacheValues = outputs[1:]
}
- return out, nil
+ return out[:steps], nil
}
func forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) {