From c51a83cfdb80933049543af41b795483a6032946 Mon Sep 17 00:00:00 2001 From: Jonas Knobloch Date: Tue, 14 Jul 2026 19:04:11 +0200 Subject: Refactor ONNX extraction --- onnx/cmd/extract/main.go | 10 +++- onnx/model.go | 146 +++++++++++++++++++++++++++++++++++++++++++++ onnx/onnx.go | 134 ----------------------------------------- research/sander/extract.go | 10 +++- 4 files changed, 164 insertions(+), 136 deletions(-) create mode 100644 onnx/model.go delete mode 100644 onnx/onnx.go diff --git a/onnx/cmd/extract/main.go b/onnx/cmd/extract/main.go index 32b8c5d..a2c281d 100644 --- a/onnx/cmd/extract/main.go +++ b/onnx/cmd/extract/main.go @@ -10,7 +10,15 @@ import ( ) func main() { - data, shape, err := onnx.ExtractInitializer(shelf.Abs("models/gpt2/model.onnx"), "transformer.wte.weight") + var model *onnx.Model + + if m, err := onnx.NewModel(shelf.Abs("models/gpt2/model.onnx")); err != nil { + log.Fatal(err) + } else { + model = m + } + + data, shape, err := model.ExtractInitializer("transformer.wte.weight") if err != nil { log.Fatal(err) diff --git a/onnx/model.go b/onnx/model.go new file mode 100644 index 0000000..cd1625c --- /dev/null +++ b/onnx/model.go @@ -0,0 +1,146 @@ +package onnx + +//go:generate bash scripts/gen.sh + +import ( + "encoding/binary" + "math" + "os" + "path/filepath" + "strconv" + + "google.golang.org/protobuf/proto" + + "go.jknobloc.com/x/onnx/internal/pb" +) + +type Model struct { + model *pb.ModelProto + name string +} + +func NewModel(name string) (*Model, error) { + var data []byte + + if d, err := os.ReadFile(name); err != nil { + return nil, err + } else { + data = d + } + + m := &Model{ + model: &pb.ModelProto{}, + name: name, + } + + if err := proto.Unmarshal(data, m.model); err != nil { + return nil, err + } + + return m, nil +} + +func (m *Model) ExtractInitializer(initializer string) ([]float32, []int, error) { + g := m.model.GetGraph() + + var result []float32 + var shape []int64 + + ok := false + + for _, init := range g.GetInitializer() { + if init.GetName() == initializer { + var raw []byte + + if init.GetDataLocation() == pb.TensorProto_EXTERNAL { + raw = readExternalData(m.name, init.GetExternalData()) + } else { + raw = init.GetRawData() + } + + result = make([]float32, len(raw)/4) + + for i := 0; i < len(result); i++ { + bits := binary.LittleEndian.Uint32(raw[i*4 : (i+1)*4]) + + result[i] = math.Float32frombits(bits) + } + + shape = init.GetDims() + + ok = true + + break + } + } + + if !ok { + panic("initializer not found") + } + + shapeInt := make([]int, len(shape)) + + for i, v := range shape { + shapeInt[i] = int(v) + } + + return result, shapeInt, nil +} + +func readExternalData(name string, entries []*pb.StringStringEntryProto) []byte { + var location string + var offset, length int64 + + for _, e := range entries { + switch e.GetKey() { + case "location": + location = e.GetValue() + case "offset": + offset, _ = strconv.ParseInt(e.GetValue(), 10, 64) + case "length": + length, _ = strconv.ParseInt(e.GetValue(), 10, 64) + } + } + + data := filepath.Join(filepath.Dir(name), location) + + var file *os.File + + if f, err := os.Open(data); err != nil { + panic(err) + } else { + file = f + } + + defer file.Close() + + if offset > 0 { + if _, err := file.Seek(offset, 0); err != nil { + panic(err) + } + } + + var buf []byte + + if length > 0 { + buf = make([]byte, length) + + if _, err := file.Read(buf); err != nil { + panic(err) + } + } else { + var err error + + buf, err = os.ReadFile(data) + + if err != nil { + panic(err) + } + + if offset > 0 { + buf = buf[offset:] + } + } + + return buf +} diff --git a/onnx/onnx.go b/onnx/onnx.go deleted file mode 100644 index b77b44e..0000000 --- a/onnx/onnx.go +++ /dev/null @@ -1,134 +0,0 @@ -package onnx - -//go:generate bash scripts/gen.sh - -import ( - "encoding/binary" - "math" - "os" - "path/filepath" - "strconv" - - "google.golang.org/protobuf/proto" - - "go.jknobloc.com/x/onnx/internal/pb" -) - -func ExtractInitializer(name string, initializer string) ([]float32, []int, error) { - var data []byte - - if d, err := os.ReadFile(name); err != nil { - return nil, nil, err - } else { - data = d - } - - model := &pb.ModelProto{} - - if err := proto.Unmarshal(data, model); err != nil { - return nil, nil, err - } - - g := model.GetGraph() - - var result []float32 - var shape []int64 - - ok := false - - for _, init := range g.GetInitializer() { - if init.GetName() == initializer { - var raw []byte - - if init.GetDataLocation() == pb.TensorProto_EXTERNAL { - raw = readExternalData(name, init.GetExternalData()) - } else { - raw = init.GetRawData() - } - - result = make([]float32, len(raw)/4) - - for i := 0; i < len(result); i++ { - bits := binary.LittleEndian.Uint32(raw[i*4 : (i+1)*4]) - - result[i] = math.Float32frombits(bits) - } - - shape = init.GetDims() - - ok = true - - break - } - } - - if !ok { - panic("initializer not found") - } - - shapeInt := make([]int, len(shape)) - - for i, v := range shape { - shapeInt[i] = int(v) - } - - return result, shapeInt, nil -} - -func readExternalData(name string, entries []*pb.StringStringEntryProto) []byte { - var location string - var offset, length int64 - - for _, e := range entries { - switch e.GetKey() { - case "location": - location = e.GetValue() - case "offset": - offset, _ = strconv.ParseInt(e.GetValue(), 10, 64) - case "length": - length, _ = strconv.ParseInt(e.GetValue(), 10, 64) - } - } - - data := filepath.Join(filepath.Dir(name), location) - - var file *os.File - - if f, err := os.Open(data); err != nil { - panic(err) - } else { - file = f - } - - defer file.Close() - - if offset > 0 { - if _, err := file.Seek(offset, 0); err != nil { - panic(err) - } - } - - var buf []byte - - if length > 0 { - buf = make([]byte, length) - - if _, err := file.Read(buf); err != nil { - panic(err) - } - } else { - var err error - - buf, err = os.ReadFile(data) - - if err != nil { - panic(err) - } - - if offset > 0 { - buf = buf[offset:] - } - } - - return buf -} diff --git a/research/sander/extract.go b/research/sander/extract.go index c72fe7f..c9a2841 100644 --- a/research/sander/extract.go +++ b/research/sander/extract.go @@ -14,7 +14,15 @@ func (e *Experiment) Extract(db *sql.DB) error { var data []float32 var shape []int - if d, s, err := onnx.ExtractInitializer(e.model, "transformer.wte.weight"); err != nil { + var model *onnx.Model + + if m, err := onnx.NewModel(e.model); err != nil { + model = m + } else { + return err + } + + if d, s, err := model.ExtractInitializer("transformer.wte.weight"); err != nil { return err } else { data = d -- cgit v1.2.3