IdentificacionIA/osnet_dinamico.py

28 lines
973 B
Python

import onnx
# 1. Rutas de archivos
modelo_estatico = "osnet_x0_25_msmt17.onnx"
modelo_dinamico = "osnet_dinamico.onnx"
print(f"Abriendo {modelo_estatico}...")
model = onnx.load(modelo_estatico)
# 2. Modificar todas las entradas (Inputs)
for input_proto in model.graph.input:
# Accedemos a la primera dimensión (índice 0), que es el Batch
dim = input_proto.type.tensor_type.shape.dim[0]
# IMPORTANTE: Debemos borrar el valor fijo (ej. 16) antes de asignar el nombre dinámico
# Esto evita conflictos en ciertas versiones de la librería ONNX
dim.ClearField('dim_value')
dim.dim_param = 'batch_size'
# 3. Modificar todas las salidas (Outputs)
for output_proto in model.graph.output:
dim = output_proto.type.tensor_type.shape.dim[0]
dim.ClearField('dim_value')
dim.dim_param = 'batch_size'
# 4. Guardar el nuevo modelo
onnx.save(model, modelo_dinamico)
print(f"¡Éxito! El nuevo modelo se ha guardado como: {modelo_dinamico}")