blob: 5a1043977e6327f3ca5788431202f3a0fcb94911 (
plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
|
#!/usr/bin/env python3
import argparse
import copy
import os
import re
import onnx
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument('--model', required=True, help='directory containing model.onnx')
parser.add_argument('--name', default='model.onnx', help='model filename within --model dir')
return parser.parse_args()
def find_layer_outputs(graph):
pattern = re.compile(r'^/transformer/h\.(\d+)/Add_1$')
layer_outputs = {}
for node in graph.node:
m = pattern.match(node.name)
if m:
layer_outputs[int(m.group(1))] = node.output[0]
return [layer_outputs[i] for i in sorted(layer_outputs)]
def make_lens_model(base_model, cut_tensor):
m = copy.deepcopy(base_model)
g = m.graph
for node in g.node:
if node.name == '/transformer/ln_f/LayerNormalization':
node.input[0] = cut_tensor
break
return m
def main():
args = parse_args()
model_dir = os.path.abspath(args.model)
base_model = onnx.load(os.path.join(model_dir, args.name), load_external_data=False)
layer_outputs = find_layer_outputs(base_model.graph)
n = len(layer_outputs)
stem = os.path.splitext(args.name)[0]
for k in range(n - 1):
cut_tensor = layer_outputs[k]
m = make_lens_model(base_model, cut_tensor)
onnx.save(m, os.path.join(model_dir, f"{stem}_lens_{k}.onnx"))
if __name__ == '__main__':
main()
|