summaryrefslogtreecommitdiff
path: root/onnx
diff options
context:
space:
mode:
Diffstat (limited to 'onnx')
-rw-r--r--onnx/cmd/extract/main.go10
-rw-r--r--onnx/model.go (renamed from onnx/onnx.go)26
2 files changed, 28 insertions, 8 deletions
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/onnx.go b/onnx/model.go
index b77b44e..cd1625c 100644
--- a/onnx/onnx.go
+++ b/onnx/model.go
@@ -14,22 +14,34 @@ import (
"go.jknobloc.com/x/onnx/internal/pb"
)
-func ExtractInitializer(name string, initializer string) ([]float32, []int, error) {
+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, nil, err
+ return nil, err
} else {
data = d
}
- model := &pb.ModelProto{}
+ m := &Model{
+ model: &pb.ModelProto{},
+ name: name,
+ }
- if err := proto.Unmarshal(data, model); err != nil {
- return nil, nil, err
+ if err := proto.Unmarshal(data, m.model); err != nil {
+ return nil, err
}
- g := model.GetGraph()
+ return m, nil
+}
+
+func (m *Model) ExtractInitializer(initializer string) ([]float32, []int, error) {
+ g := m.model.GetGraph()
var result []float32
var shape []int64
@@ -41,7 +53,7 @@ func ExtractInitializer(name string, initializer string) ([]float32, []int, erro
var raw []byte
if init.GetDataLocation() == pb.TensorProto_EXTERNAL {
- raw = readExternalData(name, init.GetExternalData())
+ raw = readExternalData(m.name, init.GetExternalData())
} else {
raw = init.GetRawData()
}