File size: 3,349 Bytes
ae2b7b0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
"""Remove the fp32 -> fp16 -> fp32 round trips that convert_float_to_float16 leaves
between two adjacent blocked (fp32) nodes.

onnxconverter_common wraps every blocked node with an input Cast(fp16->fp32) and an
output Cast(fp32->fp16), unconditionally.  When two blocked nodes are adjacent the
value is therefore squeezed through fp16 even though both endpoints are fp32.  For
this DINOv3 model that is fatal: the residual stream carries a ~1.57e5 "massive
activation" which overflows fp16's 65504 maximum and becomes inf, and the next
LayerNormalization turns inf into NaN.

This pass is a plain peephole: where Cast(to=FLOAT16) feeds Cast(to=FLOAT), the
consumer reads the original fp32 tensor instead.  Casts left with no consumers and
no role as a graph output are dropped.  Nothing else in the graph changes.
"""

import hashlib
import os
import sys
from collections import defaultdict

import onnx
from onnx import TensorProto

SRC = sys.argv[1] if len(sys.argv) > 1 else r"E:/projects/faceage-onnx/faceage-dino-fp16.onnx"
DST = sys.argv[2] if len(sys.argv) > 2 else SRC


def sha256(path):
    d = hashlib.sha256()
    with open(path, "rb") as f:
        for chunk in iter(lambda: f.read(1 << 20), b""):
            d.update(chunk)
    return d.hexdigest()


def cast_to(node):
    for a in node.attribute:
        if a.name == "to":
            return a.i
    return None


def main():
    m = onnx.load(SRC)
    g = m.graph
    graph_outputs = {o.name for o in g.output}

    producer = {o: n for n in g.node for o in n.output}
    consumers = defaultdict(list)
    for n in g.node:
        for i in n.input:
            consumers[i].append(n)

    rewired = 0
    dropped = []
    for n in list(g.node):
        if n.op_type != "Cast" or cast_to(n) != TensorProto.FLOAT:
            continue
        src = producer.get(n.input[0])
        if src is None or src.op_type != "Cast" or cast_to(src) != TensorProto.FLOAT16:
            continue
        original = src.input[0]           # the fp32 tensor before the down-cast
        target = n.output[0]
        if target in graph_outputs:
            continue
        for c in consumers[target]:
            for k, i in enumerate(c.input):
                if i == target:
                    c.input[k] = original
                    rewired += 1
        dropped.append(n)

    for n in dropped:
        g.node.remove(n)

    # Drop the now-orphaned down-casts.
    consumers = defaultdict(list)
    for n in g.node:
        for i in n.input:
            consumers[i].append(n)
    orphans = [
        n for n in g.node
        if n.op_type == "Cast" and cast_to(n) == TensorProto.FLOAT16
        and not consumers[n.output[0]] and n.output[0] not in graph_outputs
    ]
    for n in orphans:
        g.node.remove(n)

    # value_info entries for tensors that no longer exist would be stale.
    live = {o for n in g.node for o in n.output}
    keep = [v for v in g.value_info if v.name in live]
    del g.value_info[:]
    g.value_info.extend(keep)

    print("up-casts removed:", len(dropped), " consumer inputs rewired:", rewired)
    print("orphaned down-casts removed:", len(orphans))
    onnx.checker.check_model(m, full_check=False)
    onnx.save(m, DST, save_as_external_data=False)
    print("saved", DST, os.path.getsize(DST), sha256(DST))


if __name__ == "__main__":
    main()