scrfd_640_batched / fix_det_batch.py
ceyxprime's picture
Upload fix_det_batch.py with huggingface_hub
fdefdba verified
Raw
History Blame
3.29 kB
#!/usr/bin/env python3
"""
ONNX graph surgery to fix det_500m.onnx for batch-independent inference.
Root cause: 9 output Transpose nodes use perm [2,3,0,1], moving [N,C,H,W] to
[H,W,N,C]. The subsequent Reshape [-1,K] flattens everything, interleaving frames.
Fix:
1. Transpose perm [2,3,0,1] -> [0,2,3,1]: keeps batch first [N,H,W,C]
2. Reshape [-1,K] -> [0,-1,K]: preserves batch dim -> [N, anchors, K]
3. Input/output declarations updated for dynamic batch
"""
import argparse
import onnx
from onnx import numpy_helper, TensorProto, helper
import numpy as np
parser = argparse.ArgumentParser()
parser.add_argument('--model', required=True)
parser.add_argument('--out', required=True)
args = parser.parse_args()
model = onnx.load(args.model)
# 1. Make input batch dim dynamic
for inp in model.graph.input:
d0 = inp.type.tensor_type.shape.dim[0]
d0.ClearField('dim_value')
d0.dim_param = 'batch'
print(f"Input '{inp.name}': batch dim -> dynamic")
# 2. Fix the 9 output Transpose nodes: perm [2,3,0,1] -> [0,2,3,1]
output_transpose_names = {
'Transpose_154', 'Transpose_158', 'Transpose_161', # stride 8
'Transpose_179', 'Transpose_183', 'Transpose_186', # stride 16
'Transpose_204', 'Transpose_208', 'Transpose_211', # stride 32
}
patched_transposes = 0
for node in model.graph.node:
if node.name in output_transpose_names:
for attr in node.attribute:
if attr.name == 'perm':
old_perm = list(attr.ints)
attr.ints[:] = [0, 2, 3, 1]
print(f" {node.name}: perm {old_perm} -> [0, 2, 3, 1]")
patched_transposes += 1
print(f"Patched {patched_transposes} Transpose nodes")
# 3. Fix Reshape initializers: [-1, K] -> [0, -1, K]
# These are the output reshape shapes that flatten spatial+batch together
output_reshape_inits = {'441', '445', '448'}
patched_reshapes = 0
for init in list(model.graph.initializer):
if init.name not in output_reshape_inits:
continue
arr = numpy_helper.to_array(init).copy()
print(f" Reshape init '{init.name}': {arr} -> ", end='')
new_arr = np.array([0] + arr.tolist(), dtype=np.int64)
print(new_arr)
new_t = numpy_helper.from_array(new_arr, init.name)
model.graph.initializer.remove(init)
model.graph.initializer.append(new_t)
patched_reshapes += 1
print(f"Patched {patched_reshapes} Reshape initializer(s)")
# 4. Fix output shapes: 2D [anchors, K] -> 3D [batch, anchors, K]
for out in model.graph.output:
shape = out.type.tensor_type.shape
old_dims = []
for d in shape.dim:
if d.HasField('dim_param'):
old_dims.append(('param', d.dim_param))
else:
old_dims.append(('value', d.dim_value))
while len(shape.dim) > 0:
shape.dim.pop()
# Add batch dim
d = shape.dim.add()
d.dim_param = 'batch'
# Re-add original dims
for kind, val in old_dims:
d = shape.dim.add()
if kind == 'param':
d.dim_param = val
else:
d.dim_value = val
new_shape = [dd.dim_param if dd.HasField('dim_param') else dd.dim_value for dd in shape.dim]
print(f"Output '{out.name}': {old_dims} -> {new_shape}")
# Save
onnx.save(model, args.out)
print(f"\nSaved: {args.out}")