summaryrefslogtreecommitdiff
path: root/gpt2/scripts/lens.py
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()