File size: 7,694 Bytes
da9358b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
#!/usr/bin/env python
"""
Step 1: make the exported vision ONNX runnable by ONNX Runtime / plain trtexec.

The TensorRT-Edge-LLM export contains 32 `trt::ViTAttentionPlugin` nodes (custom TRT plugin, no ORT kernel).
ModelOpt ONNX PTQ calibrates with ONNX Runtime, so every plugin node is replaced by an equivalent
standard-ONNX attention sub-graph. Nothing else in the graph is touched (weights, FP16/FP32 mix, I/O
names, dtypes and shapes stay identical), so the result is still a drop-in vision encoder.

Plugin semantics (cpp/plugins/vitAttentionPlugin): q,k,v [S,H,D] fp16, cu_seqlens int32 [B+1],
ragged (block-diagonal) non-causal attention with scale 1/sqrt(D), output [S,H,D] fp16.

Replacement: seg_id[i] = #{j : i >= cu_seqlens[j]}  ->  mask[i,j] = 0 if seg_id[i]==seg_id[j] else -1e4
             out = Softmax((q*scale) @ k^T + mask) @ v   (per head)

Also: initializers that were de-duplicated by the exporter and re-used through `Identity` nodes are
copied back so that every Gemm bias is a real constant (ModelOpt/ORT quantizer wants constant bias).
"""
import argparse, math, os
from collections import Counter
import numpy as np
import onnx
import onnx_graphsurgeon as gs
from onnx import TensorProto, helper, numpy_helper

ROOT = os.environ.get("ROOT", "/data/users/logesh/Infernece_vision_Manual")
ap = argparse.ArgumentParser()
ap.add_argument("--src", default=os.environ.get("SRC_ONNX", "/data/users/logesh/TensorRT-Edge-LLM/Qwen/Qwen3-VL-2B-Instruct/onnx/visual/model.onnx"))
ap.add_argument("--dst", default=f"{ROOT}/onnx/fp16_noplugin/model.onnx")
ap.add_argument("--mask_value", type=float, default=-1e4)
ap.add_argument("--attn_fp32", action="store_true", help="mask-add + softmax in FP32 (HF eager style); default keeps the plugin's FP16")
args = ap.parse_args()

m = onnx.load(args.src, load_external_data=True)
g = m.graph
plugin_nodes = [n for n in g.node if n.op_type == "ViTAttentionPlugin"]
print(f"Found {len(plugin_nodes)} ViTAttentionPlugin nodes")

# ---------------------------------------------------------------- 1) un-alias Identity(initializer)
inits = {i.name: i for i in g.initializer}
ident = [n for n in g.node if n.op_type == "Identity" and n.input[0] in inits]
if ident:
    alias = {n.output[0]: n.input[0] for n in ident}
    new_inits = []
    for out_name, src_name in alias.items():
        t = onnx.TensorProto(); t.CopyFrom(inits[src_name]); t.name = out_name
        new_inits.append(t)
    g.initializer.extend(new_inits)
    keep = [n for n in g.node if n not in ident]
    del g.node[:]; g.node.extend(keep)
    print(f"Un-aliased {len(ident)} Identity(initializer) nodes: {list(alias.items())}")

# ---------------------------------------------------------------- 2) plugin -> standard attention
new_nodes, consts = [], []
def const(name, arr):
    consts.append(numpy_helper.from_array(arr, name)); return name

MDT = np.float32 if args.attn_fp32 else np.float16
c_zero_i64 = const("vitattn/zero_i64", np.array(0, dtype=np.int64))
c_one_i64 = const("vitattn/one_i64", np.array(1, dtype=np.int64))
c_axes0 = const("vitattn/axes0", np.array([0], dtype=np.int64))
c_axes1 = const("vitattn/axes1", np.array([1], dtype=np.int64))
c_mask0 = const("vitattn/mask_zero", np.array(0.0, dtype=MDT))
c_maskneg = const("vitattn/mask_neg", np.array(args.mask_value, dtype=MDT))

mask_cache = {}
def build_mask(cu_name, q_name):
    if cu_name in mask_cache:
        return mask_cache[cu_name]
    p = f"vitattn/mask[{cu_name}]/"
    nodes = [
        helper.make_node("Shape", [q_name], [p + "S1"], start=0, end=1),
        helper.make_node("Squeeze", [p + "S1", c_axes0], [p + "S"]),
        helper.make_node("Range", [c_zero_i64, p + "S", c_one_i64], [p + "pos"]),
        helper.make_node("Cast", [cu_name], [p + "cu64"], to=TensorProto.INT64),
        helper.make_node("Unsqueeze", [p + "pos", c_axes1], [p + "pos_col"]),        # [S,1]
        helper.make_node("Unsqueeze", [p + "cu64", c_axes0], [p + "cu_row"]),        # [1,B+1]
        helper.make_node("GreaterOrEqual", [p + "pos_col", p + "cu_row"], [p + "ge"]),
        helper.make_node("Cast", [p + "ge"], [p + "ge_i32"], to=TensorProto.INT32),
        helper.make_node("ReduceSum", [p + "ge_i32", c_axes1], [p + "seg"], keepdims=0),  # [S]
        helper.make_node("Unsqueeze", [p + "seg", c_axes1], [p + "seg_col"]),
        helper.make_node("Unsqueeze", [p + "seg", c_axes0], [p + "seg_row"]),
        helper.make_node("Equal", [p + "seg_col", p + "seg_row"], [p + "same"]),     # [S,S]
        helper.make_node("Where", [p + "same", c_mask0, c_maskneg], [p + "mask"]),
    ]
    for n in nodes:
        n.name = n.output[0]
    new_nodes.extend(nodes)
    mask_cache[cu_name] = p + "mask"
    return p + "mask"

replaced, out_nodes = 0, []
for n in g.node:
    if n.op_type != "ViTAttentionPlugin":
        out_nodes.append(n); continue
    attrs = {a.name: a.i for a in n.attribute}
    H, D = attrs["num_heads"], attrs["head_size"]
    q, k, v, cu, _carrier = n.input
    out = n.output[0]
    p = n.name + "/"
    scale_name = const(p + "scale", np.array(1.0 / math.sqrt(D), dtype=np.float16))
    mask = build_mask(cu, q)
    sub = [
        helper.make_node("Mul", [q, scale_name], [p + "q_scaled"]),
        helper.make_node("Transpose", [p + "q_scaled"], [p + "qT"], perm=[1, 0, 2]),   # [H,S,D]
        helper.make_node("Transpose", [k], [p + "kT"], perm=[1, 2, 0]),                # [H,D,S]
        helper.make_node("Transpose", [v], [p + "vT"], perm=[1, 0, 2]),                # [H,S,D]
        helper.make_node("MatMul", [p + "qT", p + "kT"], [p + "scores"]),              # [H,S,S]
    ]
    if args.attn_fp32:
        sub += [
            helper.make_node("Cast", [p + "scores"], [p + "scores32"], to=TensorProto.FLOAT),
            helper.make_node("Add", [p + "scores32", mask], [p + "scores_masked"]),
            helper.make_node("Softmax", [p + "scores_masked"], [p + "probs32"], axis=-1),
            helper.make_node("Cast", [p + "probs32"], [p + "probs"], to=TensorProto.FLOAT16),
        ]
    else:
        sub += [
            helper.make_node("Add", [p + "scores", mask], [p + "scores_masked"]),
            helper.make_node("Softmax", [p + "scores_masked"], [p + "probs"], axis=-1),
        ]
    sub += [
        helper.make_node("MatMul", [p + "probs", p + "vT"], [p + "ctx"]),              # [H,S,D]
        helper.make_node("Transpose", [p + "ctx"], [out], perm=[1, 0, 2]),             # [S,H,D]
    ]
    for s in sub:
        s.name = s.output[0]
    out_nodes.extend(sub)
    replaced += 1

# drop the now-dangling plugin carrier constants
consumed = {i for n in out_nodes for i in n.input}
out_nodes = [n for n in out_nodes if not (n.op_type == "Constant" and n.output[0].startswith("trt::ViTAttentionPlugin") and n.output[0] not in consumed)]

del g.node[:]; g.node.extend(new_nodes + out_nodes)
g.initializer.extend(consts)
keep = [o for o in m.opset_import if o.domain != "trt"]
del m.opset_import[:]; m.opset_import.extend(keep)
del g.value_info[:]

gs_graph = gs.import_onnx(m)
gs_graph.toposort().cleanup()
m = gs.export_onnx(gs_graph)
m.ir_version = 11
onnx.checker.check_model(m)
os.makedirs(os.path.dirname(args.dst), exist_ok=True)
onnx.save(m, args.dst, save_as_external_data=True, all_tensors_to_one_file=True,
          location=os.path.basename(args.dst) + ".data", size_threshold=1024)
print(f"Replaced {replaced} plugin nodes. Saved {args.dst}")
print("inputs :", [(i.name, TensorProto.DataType.Name(i.type.tensor_type.elem_type)) for i in m.graph.input])
print("outputs:", [(o.name, TensorProto.DataType.Name(o.type.tensor_type.elem_type)) for o in m.graph.output])
print(Counter(n.op_type for n in m.graph.node).most_common(14))