| |
| """ |
| 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) |
|
|
| |
| 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") |
|
|
| |
| output_transpose_names = { |
| 'Transpose_154', 'Transpose_158', 'Transpose_161', |
| 'Transpose_179', 'Transpose_183', 'Transpose_186', |
| 'Transpose_204', 'Transpose_208', 'Transpose_211', |
| } |
|
|
| 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") |
|
|
| |
| |
| 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)") |
|
|
| |
| 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() |
|
|
| |
| d = shape.dim.add() |
| d.dim_param = 'batch' |
|
|
| |
| 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}") |
|
|
| |
| onnx.save(model, args.out) |
| print(f"\nSaved: {args.out}") |
|
|