File size: 1,556 Bytes
7622ba3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
"""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")