neurogolf-solver / medal-solvers /fix_identity.py
rogermt's picture
Add fix_identity.py utility for removing Identity node waste
dbddc7d verified
Raw
History Blame
4.14 kB
"""fix_identity.py — Remove wasteful Identity nodes from ONNX models.
The Kaggle scorer counts ALL intermediate tensors except 'input' and 'output'.
If a model uses Identity as the last node (Pad → t_N → Identity → output),
the t_N tensor ([1,10,30,30] = 36KB) is counted as memory waste.
Fix: Rename the penultimate node's output to 'output' and remove Identity.
Usage:
python fix_identity.py --model optimized/task028.onnx
python fix_identity.py --dir optimized/ # Fix all .onnx in directory
"""
import onnx
import os
import argparse
import math
import numpy as np
def fix_identity(model_path, output_path=None):
"""Remove trailing Identity node from ONNX model.
Returns True if fix was applied, False if no fix needed.
"""
if output_path is None:
output_path = model_path
model = onnx.load(model_path)
# Check if last node is Identity outputting to 'output'
if len(model.graph.node) == 0:
return False
last_node = model.graph.node[-1]
if last_node.op_type != 'Identity':
return False
if 'output' not in list(last_node.output):
return False
# Get the intermediate name
intermediate_name = last_node.input[0]
# Find the node producing this intermediate and rename its output
found = False
for node in model.graph.node[:-1]:
if intermediate_name in list(node.output):
idx = list(node.output).index(intermediate_name)
node.output[idx] = 'output'
found = True
break
if not found:
return False
# Remove Identity node
del model.graph.node[-1]
# Apply strict shape inference fix
del model.graph.value_info[:]
model = onnx.shape_inference.infer_shapes(model, strict_mode=True)
onnx.save(model, output_path)
return True
def estimate_savings(model_path):
"""Estimate memory savings from removing Identity."""
model = onnx.load(model_path)
if len(model.graph.node) == 0:
return 0
last_node = model.graph.node[-1]
if last_node.op_type != 'Identity' or 'output' not in list(last_node.output):
return 0
intermediate_name = last_node.input[0]
for vi in model.graph.value_info:
if vi.name == intermediate_name:
if vi.type.HasField('tensor_type') and vi.type.tensor_type.HasField('shape'):
dims = [d.dim_value for d in vi.type.tensor_type.shape.dim]
dt = onnx.helper.tensor_dtype_to_np_dtype(vi.type.tensor_type.elem_type)
return int(np.prod(dims)) * np.dtype(dt).itemsize
# Assume [1,10,30,30] float32 if not found in value_info
return 36000
if __name__ == '__main__':
parser = argparse.ArgumentParser(description='Remove Identity node waste from ONNX models')
parser.add_argument('--model', help='Single model to fix')
parser.add_argument('--dir', help='Directory of models to fix')
parser.add_argument('--dry-run', action='store_true', help='Show savings without modifying')
args = parser.parse_args()
models = []
if args.model:
models.append(args.model)
elif args.dir:
for f in sorted(os.listdir(args.dir)):
if f.endswith('.onnx'):
models.append(os.path.join(args.dir, f))
else:
parser.print_help()
exit(1)
total_savings = 0
for model_path in models:
savings = estimate_savings(model_path)
if savings > 0:
if args.dry_run:
print(f" {model_path}: would save {savings:,} bytes")
else:
if fix_identity(model_path):
print(f" {model_path}: FIXED (saved {savings:,} bytes)")
else:
print(f" {model_path}: fix failed")
total_savings += savings
else:
print(f" {model_path}: no Identity to fix")
if total_savings > 0:
score_improvement = math.log(1 + total_savings / 50000) # rough estimate
print(f"\n Total memory savings: {total_savings:,} bytes")