final / scripts /01_replace_plugin.py
logeshsenthil13's picture
Upload folder using huggingface_hub
da9358b verified
Raw
History Blame Contribute Delete
7.69 kB
#!/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))