diff options
| author | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-05-04 15:23:34 +0200 |
|---|---|---|
| committer | Jonas Knobloch <jonas.knobloch@t-online.de> | 2026-05-04 15:28:11 +0200 |
| commit | 281e82eed24c1f871e868afbd23cb6811333c2f4 (patch) | |
| tree | 7136e246fd53848e7a33b6b7cc79203db7e3cb94 /onnx | |
| parent | 3f4748105edabbf5cefb761d1d61993d0b4d9028 (diff) | |
Add extraction from external data
Diffstat (limited to 'onnx')
| -rw-r--r-- | onnx/onnx.go | 68 |
1 files changed, 67 insertions, 1 deletions
diff --git a/onnx/onnx.go b/onnx/onnx.go index 85db619..b77b44e 100644 --- a/onnx/onnx.go +++ b/onnx/onnx.go @@ -6,6 +6,8 @@ import ( "encoding/binary" "math" "os" + "path/filepath" + "strconv" "google.golang.org/protobuf/proto" @@ -36,7 +38,13 @@ func ExtractInitializer(name string, initializer string) ([]float32, []int, erro for _, init := range g.GetInitializer() { if init.GetName() == initializer { - raw := init.GetRawData() + var raw []byte + + if init.GetDataLocation() == pb.TensorProto_EXTERNAL { + raw = readExternalData(name, init.GetExternalData()) + } else { + raw = init.GetRawData() + } result = make([]float32, len(raw)/4) @@ -66,3 +74,61 @@ func ExtractInitializer(name string, initializer string) ([]float32, []int, erro 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 +} |
