"""fp16 copy of an fp32 sherpa-onnx transducer export: onnxconverter-common float16, keep_io_types=True (inputs/outputs stay fp32, so sherpa-onnx feeds it exactly as before); metadata carried over. Shape inference runs first BY PATH (infer_shapes_path handles the >2 GB fp32 encoder), so the converter sees every intermediate type; without it a scalar Mul in pre_encode is left fp32 and the graph won't load. usage: convert_fp16.py SRC_DIR DST_DIR""" import os, shutil, sys, tempfile import onnx from onnx.shape_inference import infer_shapes_path from onnxconverter_common import float16 src, dst = sys.argv[1:3] os.makedirs(dst, exist_ok=True) for m in ("encoder", "decoder", "joiner"): inferred = f"{src}/{m}.inferred.onnx" infer_shapes_path(f"{src}/{m}.onnx", inferred) model = onnx.load(inferred) # the conv subsampling front (pre_encode, ~0.1 % of the FLOPs) stays fp32: the converter mis-types its # length-mask Cast/Mul otherwise keep32 = [n.name for n in model.graph.node if n.name.startswith("/pre_encode/")] m16 = float16.convert_float_to_float16(model, keep_io_types=True, disable_shape_infer=True, node_block_list=keep32) onnx.save(m16, f"{dst}/{m}.fp16.onnx") os.remove(inferred) shutil.copy(f"{src}/tokens.txt", f"{dst}/tokens.txt") print("ok", dst)