From 7073b124c5eb31169442ef18d1627b7ed7401280 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Thu, 13 Nov 2025 23:04:57 +0100 Subject: Scaffold byte-pair correction --- gpt2/.gitignore | 1 + gpt2/cmd/gpt2/main.go | 18 ++-- gpt2/main.go | 232 -------------------------------------------- gpt2/model.go | 264 ++++++++++++++++++++++++++++++++++++++++++++++++++ gpt2/scripts/conv.py | 2 +- 5 files changed, 275 insertions(+), 242 deletions(-) create mode 100644 gpt2/.gitignore delete mode 100644 gpt2/main.go create mode 100644 gpt2/model.go (limited to 'gpt2') diff --git a/gpt2/.gitignore b/gpt2/.gitignore new file mode 100644 index 0000000..8c6790b --- /dev/null +++ b/gpt2/.gitignore @@ -0,0 +1 @@ +/models \ No newline at end of file diff --git a/gpt2/cmd/gpt2/main.go b/gpt2/cmd/gpt2/main.go index ff84f42..96dabeb 100644 --- a/gpt2/cmd/gpt2/main.go +++ b/gpt2/cmd/gpt2/main.go @@ -4,24 +4,24 @@ import ( "fmt" "gpt2" "log" - - ort "github.com/yalue/onnxruntime_go" ) func main() { - ort.SetSharedLibraryPath("lib/onnxruntime-osx-arm64-1.22.0/lib/libonnxruntime.1.22.0.dylib") + prompt := []int64{464, 2068, 7586, 21831} + + m := gpt2.NewModel("models/base/model.onnx") - if err := ort.InitializeEnvironment(); err != nil { + if err := m.Init(); err != nil { log.Fatal(err) } - defer ort.DestroyEnvironment() - - prompt := []int64{464, 2068, 7586, 21831} - - if out, err := gpt2.Generate("scripts/onnx-gpt2/model.onnx", prompt, 5, nil); err != nil { + if out, err := m.Generate(prompt, 5, nil); err != nil { log.Fatal(err) } else { fmt.Printf("\n%v\n", out) } + + if err := m.Destroy(); err != nil { + log.Fatal(err) + } } diff --git a/gpt2/main.go b/gpt2/main.go deleted file mode 100644 index ee9e19f..0000000 --- a/gpt2/main.go +++ /dev/null @@ -1,232 +0,0 @@ -package gpt2 - -import ( - "errors" - "fmt" - "log" - "math" - "sort" - - ort "github.com/yalue/onnxruntime_go" -) - -const ( - vocabSize = 50257 - nLayers = 12 - nHeads = 12 - headDim = 64 -) - -func Generate(model string, prompt []int64, steps int64, logits *[][]float32) ([]int64, error) { - if len(prompt) == 0 { - return nil, errors.New("empty prompt") - } - - context := int64(len(prompt)) - - cacheNames, cacheValues := emptyCache() - - token := prompt[0] - - out := make([]int64, 0, steps+1) - - for step := range context + steps { - _, _, outputs, err := forward(model, token, step, cacheNames, cacheValues) - - if err != nil { - return nil, err - } - - l := outputs[0].(*ort.Tensor[float32]).GetData() - - if logits != nil { - *logits = append(*logits, l) - } - - idx, p := topK(softmax(l), 5) - - fmt.Printf("\n%d\n\n", token) - - for i, t := range idx { - fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t) - } - - 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[:steps], nil -} - -func forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) { - inputNames, inputs, _ := initInputs(token, position) - outputNames, outputs, logits, _ := initOutputs(position) - - inputNames = append(inputNames, cacheNames...) - inputs = append(inputs, cacheValues...) - - session, err := ort.NewAdvancedSession( - model, - inputNames, - outputNames, - inputs, - outputs, - nil, - ) - - if err != nil { - log.Fatal(err) - } - - defer session.Destroy() - - if err := session.Run(); err != nil { - log.Fatal(err) - } - - return logits, outputNames, outputs, nil -} - -func emptyCache() ([]string, []ort.Value) { - names := make([]string, 0, 2*nLayers) - values := make([]ort.Value, 0, 2*nLayers) - shape := []int64{1, int64(nHeads), 0, int64(headDim)} - - for i := range nLayers { - kName := fmt.Sprintf("past_key_values.%d.key", i) - vName := fmt.Sprintf("past_key_values.%d.value", i) - - kTensor, _ := ort.NewEmptyTensor[float32](shape) - vTensor, _ := ort.NewEmptyTensor[float32](shape) - - names = append(names, kName, vName) - values = append(values, ort.Value(kTensor), ort.Value(vTensor)) - } - - return names, values -} - -func initInputs(token, position int64) ([]string, []ort.Value, error) { - inputNames := []string{"input_ids", "position_ids", "attention_mask"} - - var tokens *ort.Tensor[int64] - var positions *ort.Tensor[int64] - var attentionMask *ort.Tensor[int64] - - if t, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{token}); err != nil { - return nil, nil, err - } else { - tokens = t - } - - if p, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{position}); err != nil { - return nil, nil, err - } else { - positions = p - } - - maskData := make([]int64, position+1) - maskShape := []int64{1, position + 1} - - for i := range maskData { - maskData[i] = 1 - } - - if m, err := ort.NewTensor[int64](maskShape, maskData); err != nil { - return nil, nil, err - } else { - attentionMask = m - } - - inputs := []ort.Value{ort.Value(tokens), ort.Value(positions), ort.Value(attentionMask)} - - return inputNames, inputs, nil -} - -func initOutputs(position int64) ([]string, []ort.Value, *ort.Tensor[float32], error) { - outputNames := make([]string, 0, 1+2*nLayers) - outputValues := make([]ort.Value, 0, 1+2*nLayers) - - logits, err := ort.NewEmptyTensor[float32]([]int64{1, 1, int64(vocabSize)}) - - if err != nil { - return nil, nil, nil, err - } - - outputNames = append(outputNames, "logits") - outputValues = append(outputValues, ort.Value(logits)) - - shape := []int64{1, int64(nHeads), position + 1, int64(headDim)} - - for i := range nLayers { - kName := fmt.Sprintf("present.%d.key", i) - vName := fmt.Sprintf("present.%d.value", i) - - kTensor, _ := ort.NewEmptyTensor[float32](shape) - vTensor, _ := ort.NewEmptyTensor[float32](shape) - - outputNames = append(outputNames, kName, vName) - outputValues = append(outputValues, ort.Value(kTensor), ort.Value(vTensor)) - } - - return outputNames, outputValues, logits, nil -} - -func softmax(logits []float32) []float32 { - m := logits[0] - - for _, v := range logits { - if v > m { - m = v - } - } - - s := float32(0.0) - r := make([]float32, len(logits)) - - for i, v := range logits { - e := float32(math.Exp(float64(v - m))) - - r[i] = e - s += e - } - - for i := range r { - r[i] /= s - } - - return r -} - -func topK(p []float32, k int) ([]int, []float32) { - n := len(p) - - if k > n { - k = n - } - - idx := make([]int, n) - - for i := range idx { - idx[i] = i - } - - sort.Slice(idx, func(i, j int) bool { - return p[idx[i]] > p[idx[j]] - }) - - topIdx := idx[:k] - topP := make([]float32, k) - - for i := 0; i < k; i++ { - topP[i] = p[topIdx[i]] - } - - return topIdx, topP -} diff --git a/gpt2/model.go b/gpt2/model.go new file mode 100644 index 0000000..1b7a9d1 --- /dev/null +++ b/gpt2/model.go @@ -0,0 +1,264 @@ +package gpt2 + +import ( + "errors" + "fmt" + _ "llm" + "log" + "math" + "os" + "sort" + + ort "github.com/yalue/onnxruntime_go" +) + +const ( + vocabSize = 50257 + nLayers = 12 + nHeads = 12 + headDim = 64 +) + +type Model struct { + name string +} + +func NewModel(name string) *Model { + return &Model{ + name: name, + } +} + +func (m *Model) SharedLibraryPath() string { + p, ok := os.LookupEnv("ONNXRUNTIME_SHARED_LIBRARY_PATH") + + if !ok { + // TODO embed runtime binaries + } + + return p +} + +func (m *Model) Init() error { + ort.SetSharedLibraryPath(m.SharedLibraryPath()) + + return ort.InitializeEnvironment() +} + +func (m *Model) Destroy() error { + return ort.DestroyEnvironment() +} + +func (m *Model) Generate(prompt []int64, steps int64, logits *[][]float32) ([]int64, error) { + if len(prompt) == 0 { + return nil, errors.New("empty prompt") + } + + context := int64(len(prompt)) + + cacheNames, cacheValues := emptyCache() + + token := prompt[0] + + out := make([]int64, 0, steps+1) + + for step := range context + steps { + _, _, outputs, err := forward(m.name, token, step, cacheNames, cacheValues) + + if err != nil { + return nil, err + } + + l := outputs[0].(*ort.Tensor[float32]).GetData() + + if logits != nil { + *logits = append(*logits, l) + } + + idx, _ := topK(softmax(l), 5) + + // fmt.Printf("\n%d\n\n", token) + + // for i, t := range idx { + // fmt.Printf("%.4f %.4f [%d]\n", l[t], p[i], t) + // } + + 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[:steps], nil +} + +func forward(model string, token int64, position int64, cacheNames []string, cacheValues []ort.Value) (*ort.Tensor[float32], []string, []ort.Value, error) { + inputNames, inputs, _ := initInputs(token, position) + outputNames, outputs, logits, _ := initOutputs(position) + + inputNames = append(inputNames, cacheNames...) + inputs = append(inputs, cacheValues...) + + session, err := ort.NewAdvancedSession( + model, + inputNames, + outputNames, + inputs, + outputs, + nil, + ) + + if err != nil { + log.Fatal(err) + } + + defer session.Destroy() + + if err := session.Run(); err != nil { + log.Fatal(err) + } + + return logits, outputNames, outputs, nil +} + +func emptyCache() ([]string, []ort.Value) { + names := make([]string, 0, 2*nLayers) + values := make([]ort.Value, 0, 2*nLayers) + shape := []int64{1, int64(nHeads), 0, int64(headDim)} + + for i := range nLayers { + kName := fmt.Sprintf("past_key_values.%d.key", i) + vName := fmt.Sprintf("past_key_values.%d.value", i) + + kTensor, _ := ort.NewEmptyTensor[float32](shape) + vTensor, _ := ort.NewEmptyTensor[float32](shape) + + names = append(names, kName, vName) + values = append(values, ort.Value(kTensor), ort.Value(vTensor)) + } + + return names, values +} + +func initInputs(token, position int64) ([]string, []ort.Value, error) { + inputNames := []string{"input_ids", "position_ids", "attention_mask"} + + var tokens *ort.Tensor[int64] + var positions *ort.Tensor[int64] + var attentionMask *ort.Tensor[int64] + + if t, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{token}); err != nil { + return nil, nil, err + } else { + tokens = t + } + + if p, err := ort.NewTensor[int64]([]int64{1, 1}, []int64{position}); err != nil { + return nil, nil, err + } else { + positions = p + } + + maskData := make([]int64, position+1) + maskShape := []int64{1, position + 1} + + for i := range maskData { + maskData[i] = 1 + } + + if m, err := ort.NewTensor[int64](maskShape, maskData); err != nil { + return nil, nil, err + } else { + attentionMask = m + } + + inputs := []ort.Value{ort.Value(tokens), ort.Value(positions), ort.Value(attentionMask)} + + return inputNames, inputs, nil +} + +func initOutputs(position int64) ([]string, []ort.Value, *ort.Tensor[float32], error) { + outputNames := make([]string, 0, 1+2*nLayers) + outputValues := make([]ort.Value, 0, 1+2*nLayers) + + logits, err := ort.NewEmptyTensor[float32]([]int64{1, 1, int64(vocabSize)}) + + if err != nil { + return nil, nil, nil, err + } + + outputNames = append(outputNames, "logits") + outputValues = append(outputValues, ort.Value(logits)) + + shape := []int64{1, int64(nHeads), position + 1, int64(headDim)} + + for i := range nLayers { + kName := fmt.Sprintf("present.%d.key", i) + vName := fmt.Sprintf("present.%d.value", i) + + kTensor, _ := ort.NewEmptyTensor[float32](shape) + vTensor, _ := ort.NewEmptyTensor[float32](shape) + + outputNames = append(outputNames, kName, vName) + outputValues = append(outputValues, ort.Value(kTensor), ort.Value(vTensor)) + } + + return outputNames, outputValues, logits, nil +} + +func softmax(logits []float32) []float32 { + m := logits[0] + + for _, v := range logits { + if v > m { + m = v + } + } + + s := float32(0.0) + r := make([]float32, len(logits)) + + for i, v := range logits { + e := float32(math.Exp(float64(v - m))) + + r[i] = e + s += e + } + + for i := range r { + r[i] /= s + } + + return r +} + +func topK(p []float32, k int) ([]int, []float32) { + n := len(p) + + if k > n { + k = n + } + + idx := make([]int, n) + + for i := range idx { + idx[i] = i + } + + sort.Slice(idx, func(i, j int) bool { + return p[idx[i]] > p[idx[j]] + }) + + topIdx := idx[:k] + topP := make([]float32, k) + + for i := 0; i < k; i++ { + topP[i] = p[topIdx[i]] + } + + return topIdx, topP +} diff --git a/gpt2/scripts/conv.py b/gpt2/scripts/conv.py index cf0cbb8..e3773e2 100644 --- a/gpt2/scripts/conv.py +++ b/gpt2/scripts/conv.py @@ -12,4 +12,4 @@ from optimum.onnxruntime import ORTModelForCausalLM model_id = "gpt2" model = ORTModelForCausalLM.from_pretrained(model_id, export=True, use_cache=True) -model.save_pretrained("onnx-gpt2") \ No newline at end of file +model.save_pretrained("../models/base") \ No newline at end of file -- cgit v1.3.1