import onnx, onnxruntime as ort, re from onnxconverter_common import float16 def convert(src_onnx, dst_onnx, dst_data, max_rounds=8): inferred = src_onnx.replace(".onnx", "-inferred.onnx") onnx.shape_inference.infer_shapes_path(src_onnx, inferred) # >2GB-safe shape inference base = onnx.load(inferred, load_external_data=False) block = set(n.name for n in base.graph.node if "pre_encode" in n.name) # round 0: keep front-end fp32 for rnd in range(max_rounds): m = onnx.load(inferred, load_external_data=True) m16 = float16.convert_float_to_float16(m, keep_io_types=True, disable_shape_infer=True, node_block_list=sorted(block)) del m16.graph.value_info[:] # let ORT re-infer consistent types onnx.save_model(m16, dst_onnx, save_as_external_data=True, all_tensors_to_one_file=True, location=dst_data, size_threshold=1024) try: ort.InferenceSession(dst_onnx, providers=["CPUExecutionProvider"]) except Exception as e: hit = re.search(r"node \(([^)]+)\)", str(e)) # which node did ORT reject? if not hit: raise module = hit.group(1).rsplit("/", 1)[0] # keep that whole submodule fp32, retry add = set(n.name for n in m16.graph.node if n.name == module or n.name.startswith(module + "/")) if add <= block: raise print("round %d: ORT rejected %s -> also keeping %s fp32, retrying" % (rnd, hit.group(1), module)) block |= add continue n_fp16 = sum(1 for i in m16.graph.initializer if i.data_type == onnx.TensorProto.FLOAT16) print("OK + ORT-LOADS: %s | nodes %d | kept_fp32_nodes %d | fp16_initializers %d" % (dst_onnx, len(m16.graph.node), len(block), n_fp16)) return raise RuntimeError("%s: not ORT-loadable after %d rounds" % (dst_onnx, max_rounds)) convert("encoder-model.onnx", "encoder_model.onnx", "encoder_model.onnx.data") convert("decoder_joint-model.onnx", "decoder_model.onnx", "decoder_model.onnx.data")