laya-onnx / scripts /quant.py
techtheist's picture
Laya ONNX exports: en int4/int8, multilingual int8, dynamic sequence length
7622ba3 verified
Raw History Blame Contribute Delete
1.56 kB
"""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")