"""int8 (dynamic, MatMulInteger) and int4 (MatMulNBits + GatherBlockQuantized) from an fp32 export.""" import sys, os, time import onnx src, outdir = sys.argv[1], sys.argv[2] which = sys.argv[3:] or ["int8", "int4"] clean = os.path.join(outdir, "clean.onnx") if not os.path.exists(clean): m0 = onnx.load(src) del m0.graph.value_info[:] onnx.save(m0, clean, save_as_external_data=True, location="clean.onnx.data") del m0 src = clean if "int8" in which: from onnxruntime.quantization import quantize_dynamic, QuantType t = time.time() quantize_dynamic(src, os.path.join(outdir, "model_int8.onnx"), weight_type=QuantType.QInt8, per_channel=True, op_types_to_quantize=["MatMul", "Gemm"], extra_options={"MatMulConstBOnly": True}) print("int8 %.0fs" % (time.time() - t)) if "int4" in which: from onnxruntime.quantization.matmul_nbits_quantizer import MatMulNBitsQuantizer, DefaultWeightOnlyQuantConfig t = time.time() m = onnx.load(src) cfg = DefaultWeightOnlyQuantConfig(block_size=32, is_symmetric=True, accuracy_level=4, op_types_to_quantize=("MatMul", "Gather"), quant_axes=(("MatMul", 0), ("Gather", 1))) q = MatMulNBitsQuantizer(m, algo_config=cfg) q.process() q.model.save_model_to_file(os.path.join(outdir, "model_int4.onnx"), use_external_data_format=False) print("int4 %.0fs" % (time.time() - t)) for f in sorted(os.listdir(outdir)): print(f, os.path.getsize(os.path.join(outdir, f)) // 2**20, "MB")