codekingpro/portable-devtools
115k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6
7from __future__ import annotations
8
9import argparse
10import copy
11import logging
12import os
13
14import ml_dtypes
15import numpy as np
16import numpy.typing as npt
17import onnx
18import onnx_ir as ir
19from onnx.onnx_pb import GraphProto, ModelProto, NodeProto, TensorProto
20
21from onnxruntime.capi._pybind_state import (
22 quantize_matmul_2bits,
23 quantize_matmul_4bits,
24 quantize_matmul_8bits,
25 quantize_qdq_matmul_4bits,
26)
27
28from .calibrate import CalibrationDataReader
29from .neural_compressor import gptq_quantize, rtn_quantize
30from .onnx_model import ONNXModel
31from .quant_utils import QuantFormat, attribute_to_kwarg
32
33logging.basicConfig(format="%(asctime)s %(name)s [%(levelname)s] - %(message)s", level=logging.INFO)
34logger = logging.getLogger(__name__)
35
36
37class WeightOnlyQuantConfig:
38 def __init__(
39 self,
40 algorithm: str,
41 quant_format: QuantFormat,
42 op_types_to_quantize: tuple[str, ...] | None = None,
43 quant_axes: tuple[tuple[str, int], ...] | None = None,
44 customized_weight_config: dict | None = None,
45 ):
46 """This is the Base class for Weight Only blockwise quantization Configuration.
47
48 Args:
49 algorithm:
50 weight only quantize algorithm name.
51 quant_format: QuantFormat{QOperator, QDQ}.
52 QOperator format quantizes the model with quantized operators directly.
53 QDQ format quantize the model by inserting QuantizeLinear/DeQuantizeLinear on the tensor.
54 op_types_to_quantize (optional):
55 set of operator types to quantize. Default {MatMul}
56 quant_axes (dict[str, int], optional):
57 op:axis, which axis to quantize for an op. Default {MatMul: 0, Gather: 1}
58 customized_weight_config:
59 customized weight config for nodes if needed. It is dictionary with node name as key,
60 and the value is a dict of customized config.
61 """
62 self.algorithm = algorithm
63 self.quant_format = quant_format
64 self.op_types_to_quantize = set(op_types_to_quantize) if op_types_to_quantize else {"MatMul"}
65 self.quant_axes = dict(quant_axes) if quant_axes else {"MatMul": 0, "Gather": 1}
66 self.customized_weight_config = customized_weight_config
67
68
69class RTNWeightOnlyQuantConfig(WeightOnlyQuantConfig):
70 def __init__(
71 self,
72 ratios=None,
73 quant_format=QuantFormat.QOperator,
74 op_types_to_quantize: tuple[str, ...] | None = None,
75 customized_weight_config: dict | None = None,
76 ):
77 """
78 This is a class for round-to-nearest (RTN) algorithm Weight Only Quant Configuration.
79 RTN is the most straightforward way to quantize weight using scale maps.
80
81 Args:
82 ratios:
83 percentile of clip. Defaults to {}.
84 quant_format (QuantFormat{QOperator, QDQ}, optional):
85 QOperator format quantizes the model with quantized operators directly.
86 QDQ format quantize the model by inserting QuantizeLinear/DeQuantizeLinear on the tensor.
87 Defaults to QuantFormat.QOperator.
88 op_types_to_quantize (optional):
89 set of operator types to quantize.
90 customized_weight_config:
91 customized weight config for nodes if needed. It is dictionary with node name as key,
92 and the value is a dict of customized config.
93 """
94 assert quant_format == QuantFormat.QOperator, "RTN only supports QOperator format"
95
96 if ratios is None:
97 ratios = {}
98 super().__init__(
99 algorithm="RTN",
100 quant_format=quant_format,
101 op_types_to_quantize=op_types_to_quantize,
102 customized_weight_config=customized_weight_config,
103 )
104 self.ratios = ratios
105
106
107class KQuantWeightOnlyQuantConfig(WeightOnlyQuantConfig):
108 def __init__(
109 self,
110 ratios=None,
111 quant_format=QuantFormat.QOperator,
112 op_types_to_quantize: tuple[str, ...] | None = None,
113 customized_weight_config: dict | None = None,
114 ):
115 """
116 This is a class for k-quant algorithm Weight Only Quant Configuration.
117
118 Args:
119 ratios:
120 percentile of clip. Defaults to {}.
121 quant_format (QuantFormat{QOperator, QDQ}, optional):
122 QOperator format quantizes the model with quantized operators directly.
123 QDQ format quantize the model by inserting QuantizeLinear/DeQuantizeLinear on the tensor.
124 Defaults to QuantFormat.QOperator.
125 op_types_to_quantize (optional):
126 set of operator types to quantize.
127 """
128 assert quant_format == QuantFormat.QOperator, "k-quant only supports QOperator format"
129
130 if ratios is None:
131 ratios = {}
132 super().__init__(
133 algorithm="k_quant",
134 quant_format=quant_format,
135 op_types_to_quantize=op_types_to_quantize,
136 customized_weight_config=customized_weight_config,
137 )
138 self.ratios = ratios
139
140
141class GPTQWeightOnlyQuantConfig(WeightOnlyQuantConfig):
142 def __init__(
143 self,
144 calibration_data_reader: CalibrationDataReader | None = None,
145 percdamp=0.01,
146 block_size=128,
147 actorder=False,
148 mse=False,
149 perchannel=True,
150 quant_format=QuantFormat.QOperator,
151 op_types_to_quantize: tuple[str, ...] | None = None,
152 ):
153 """
154 This is a class for GPTQ algorithm Weight Only Quant Configuration.
155 GPTQ algorithm provides more accurate quantization but requires more computational resources.
156
157 Args:
158 calibration_data_reader:
159 a calibration data reader. It enumerates calibration data and generates inputs for the original model.
160 percdamp:
161 percent of the average Hessian diagonal to use for dampening.
162 block_size (int, optional):
163 channel number in one block to execute a GPTQ quantization iteration.
164 actorder (bool, optional):
165 whether rearrange Hessian matrix considering the diag's value.
166 mse (bool, optional):
167 whether get scale and zero point with mse error.
168 perchannel (bool, optional):
169 whether quantize weight per-channel.
170 quant_format (QuantFormat{QOperator, QDQ}, optional):
171 QOperator format quantizes the model with quantized operators directly.
172 QDQ format quantize the model by inserting QuantizeLinear/DeQuantizeLinear on the tensor.
173 Defaults to QuantFormat.QOperator.
174 op_types_to_quantize (optional):
175 set of operator types to quantize.
176 """
177 assert quant_format == QuantFormat.QOperator, "GPTQ only supports QOperator format"
178
179 super().__init__(
180 algorithm="GPTQ",
181 quant_format=quant_format,
182 op_types_to_quantize=op_types_to_quantize,
183 )
184 self.calibration_data_reader = calibration_data_reader
185 self.percdamp = percdamp
186 self.block_size = block_size
187 self.actorder = actorder
188 self.mse = mse
189 self.perchannel = perchannel
190
191
192class HQQWeightOnlyQuantConfig(WeightOnlyQuantConfig):
193 def __init__(
194 self,
195 block_size=128,
196 bits=4,
197 axis=1,
198 quant_format=QuantFormat.QOperator,
199 op_types_to_quantize: tuple[str, ...] | None = None,
200 quant_axes: tuple[tuple[str, int], ...] | None = None,
201 ):
202 """
203 This is a class for HQQ algorithm Weight Only Quant Configuration.
204 HQQ algorithm quant weight without needing calibrate data.
205
206 Args:
207 block_size (int, optional):
208 channel number in one block to execute a HQQ quantization iteration.
209 bits (int, optional):
210 how many bits to represent weight.
211 axis (int, optional):
212 0 or 1. which axis to quantize. https://arxiv.org/pdf/2309.15531.pdf
213 quant_format (QuantFormat{QOperator, QDQ}, optional):
214 QOperator format quantizes the model with quantized operators directly.
215 QDQ format quantize the model by inserting QuantizeLinear/DeQuantizeLinear on the tensor.
216 Defaults to QuantFormat.QOperator.
217 op_types_to_quantize (optional):
218 set of operator types to quantize.
219 quant_axes (dict[str, int], optional):
220 op:axis, which axis to quantize for an op. Default {MatMul: 0, Gather: 1}
221 """
222 assert quant_format == QuantFormat.QOperator, "HQQ only supports QOperator format"
223
224 super().__init__(
225 algorithm="HQQ",
226 quant_format=quant_format,
227 op_types_to_quantize=op_types_to_quantize,
228 quant_axes=quant_axes,
229 )
230 self.block_size = block_size
231 self.bits = bits
232 self.axis = axis
233
234
235class DefaultWeightOnlyQuantConfig(WeightOnlyQuantConfig):
236 def __init__(
237 self,
238 block_size: int = 128,
239 is_symmetric: bool = False,
240 accuracy_level: int | None = None,
241 quant_format=QuantFormat.QOperator,
242 op_types_to_quantize: tuple[str, ...] | None = None,
243 quant_axes: tuple[tuple[str, int], ...] | None = None,
244 bits: int = 4,
245 channel_wised_quantize: bool = False,
246 ):
247 """
248 This is a class for weight only affine quantization configuration.
249
250 Args:
251 block_size (int, optional):
252 channel number in one block to execute an affine quantization iteration.
253 is_symmetric (bool, optional):
254 whether quantize weight symmetrically.
255 accuracy_level (int, optional):
256 Accuracy level of the 4-bit quantized MatMul computation.
257 Refer to the MatMulNBits contrib op's 'accuracy_level' attribute for details.
258 (https://github.com/microsoft/onnxruntime/blob/main/docs/ContribOperators.md#commicrosoftmatmulnbits)
259 quant_format (QuantFormat{QOperator, QDQ}, optional):
260 QOperator format quantizes the model with quantized operators directly.
261 QDQ format quantize the model by inserting QuantizeLinear/DeQuantizeLinear on the tensor.
262 Defaults to QuantFormat.QOperator.
263 op_types_to_quantize (optional):
264 set of operator types to quantize.
265 quant_axes (dict[str, int], optional):
266 op:axis, which axis to quantize for an op. Default {MatMul: 0, Gather: 1}
267 bits (int, optional):
268 number of bits per element after quantization. Default 4.
269 """
270 super().__init__(
271 algorithm="DEFAULT",
272 quant_format=quant_format,
273 op_types_to_quantize=op_types_to_quantize,
274 quant_axes=quant_axes,
275 )
276 self.block_size = block_size
277 self.is_symmetric = is_symmetric
278 self.bits = bits
279 self.accuracy_level = accuracy_level
280 self.channel_wised_quantize = channel_wised_quantize
281 if channel_wised_quantize and quant_format == QuantFormat.QOperator:
282 raise NotImplementedError("QuantFormat.QOperator is not supported channel_wised_quantize yet")
283
284
285class NVAWQWeightOnlyQuantConfig(WeightOnlyQuantConfig):
286 def __init__(
287 self,
288 tokenizer_dir,
289 dataset_name="cnn",
290 cache_dir="./cache",
291 calibration_method="awq_lite",
292 ):
293 """
294 Configuration for the nvidia_awq quantization method.
295
296 Args:
297 tokenizer_dir (str): pathof the tokenizer dir.
298 dataset_name (str): Name of the dataset.
299 cache_dir (str): Directory for caching.
300 calibration_method (str): calib method for nvidia_awq.
301 """
302 # Import torch and DataLoader
303 try:
304 import torch # noqa: PLC0415
305 from torch.utils.data import DataLoader # noqa: PLC0415
306
307 self.torch = torch
308 self.DataLoader = DataLoader
309 except ImportError:
310 print(
311 "Error: The 'torch' library is required but not installed. Please install it using 'pip install torch'."
312 )
313 raise ImportError("torch is not installed. Exiting.") from None
314
315 # Import datasets
316 try:
317 from datasets import load_dataset # noqa: PLC0415
318
319 self.load_dataset = load_dataset
320 except ImportError:
321 print(
322 "Error: The 'datasets' library is required but not installed. Please install it using 'pip install datasets'."
323 )
324 raise ImportError("datasets is not installed. Exiting.") from None
325
326 # Import transformers
327 try:
328 from transformers import AutoConfig, AutoTokenizer # noqa: PLC0415
329
330 self.AutoConfig = AutoConfig
331 self.AutoTokenizer = AutoTokenizer
332 except ImportError:
333 print(
334 "Error: The 'transformers' library is required but not installed. Please install it using 'pip install transformers'."
335 )
336 raise ImportError("transformers is not installed. Exiting.") from None
337
338 super().__init__(
339 algorithm="nvidia_awq",
340 quant_format=QuantFormat.QDQ,
341 op_types_to_quantize=None, # Assuming op_types_to_quantize is handled elsewhere
342 quant_axes=None, # Assuming quant_axes is handled elsewhere
343 )
344
345 # Determine the device
346 device = self.torch.device("cuda" if self.torch.cuda.is_available() else "cpu")
347
348 calib_inputs = self.get_calib_inputs(
349 dataset_name=dataset_name,
350 model_name=tokenizer_dir,
351 cache_dir=cache_dir,
352 calib_size=32,
353 batch_size=1,
354 block_size=512,
355 device=device,
356 use_fp16=True,
357 use_buffer_share=False,
358 add_past_kv_inputs=True,
359 max_calib_rows_to_load=128,
360 add_position_ids=True,
361 )
362
363 self.calibration_data_reader = calib_inputs
364 self.calibration_method = calibration_method
365
366 def make_model_input(
367 self,
368 config,
369 input_ids_arg,
370 attention_mask_arg,
371 add_past_kv_inputs,
372 device,
373 use_fp16,
374 use_buffer_share,
375 add_position_ids,
376 ):
377 # Access torch from the instance variable
378 torch = self.torch
379
380 input_ids = input_ids_arg
381 attention_mask = attention_mask_arg
382
383 if isinstance(input_ids_arg, list):
384 input_ids = torch.tensor(input_ids_arg, device=device, dtype=torch.int64)
385 attention_mask = torch.tensor(attention_mask_arg, device=device, dtype=torch.int64)
386
387 inputs = {
388 "input_ids": input_ids.contiguous(),
389 "attention_mask": attention_mask.contiguous(),
390 }
391
392 if add_position_ids:
393 position_ids = attention_mask.long().cumsum(-1) - 1
394 position_ids.masked_fill_(attention_mask == 0, 1)
395 inputs["position_ids"] = position_ids.contiguous()
396
397 if add_past_kv_inputs:
398 torch_dtype = torch.float16 if use_fp16 else torch.float32
399 batch_size, sequence_length = input_ids.shape
400 max_sequence_length = config.max_position_embeddings
401 num_heads, head_size = (
402 config.num_key_value_heads,
403 config.hidden_size // config.num_attention_heads,
404 )
405 for i in range(config.num_hidden_layers):
406 past_key = torch.zeros(
407 batch_size,
408 num_heads,
409 max_sequence_length if use_buffer_share else 0,
410 head_size,
411 device=device,
412 dtype=torch_dtype,
413 )
414 past_value = torch.zeros(
415 batch_size,
416 num_heads,
417 max_sequence_length if use_buffer_share else 0,
418 head_size,
419 device=device,
420 dtype=torch_dtype,
421 )
422 inputs.update(
423 {
424 f"past_key_values.{i}.key": past_key.contiguous(),
425 f"past_key_values.{i}.value": past_value.contiguous(),
426 }
427 )
428
429 return inputs
430
431 def get_calib_inputs(
432 self,
433 dataset_name,
434 model_name,
435 cache_dir,
436 calib_size,
437 batch_size,
438 block_size,
439 device,
440 use_fp16,
441 use_buffer_share,
442 add_past_kv_inputs,
443 max_calib_rows_to_load,
444 add_position_ids,
445 ):
446 # Access transformers and datasets from the instance variables
447 auto_config = self.AutoConfig
448 auto_tokenizer = self.AutoTokenizer
449 load_dataset = self.load_dataset
450
451 config = auto_config.from_pretrained(
452 model_name, use_auth_token=True, cache_dir=cache_dir, trust_remote_code=True
453 )
454 tokenizer = auto_tokenizer.from_pretrained(
455 model_name, use_auth_token=True, cache_dir=cache_dir, trust_remote_code=True
456 )
457 tokenizer.add_special_tokens({"pad_token": "[PAD]"})
458 tokenizer.pad_token = tokenizer.eos_token
459
460 assert calib_size <= max_calib_rows_to_load, "calib size should be no more than max_calib_rows_to_load"
461
462 if "cnn" in dataset_name:
463 dataset2 = load_dataset("cnn_dailymail", name="3.0.0", split="train").select(range(max_calib_rows_to_load))
464 column = "article"
465 elif "pile" in dataset_name:
466 dataset2 = load_dataset("mit-han-lab/pile-val-backup", split="validation")
467 column = "text"
468 else:
469 raise ValueError(f'dataset "{dataset_name}" not supported')
470
471 dataset2 = dataset2[column][:calib_size]
472 batch_encoded = tokenizer.batch_encode_plus(
473 dataset2, return_tensors="pt", padding=True, truncation=True, max_length=block_size
474 )
475 batch_encoded = batch_encoded.to(device)
476 batch_encoded_input_ids = batch_encoded["input_ids"]
477 batch_encoded_attention_mask = batch_encoded["attention_mask"]
478
479 # Access DataLoader from the instance variable
480 data_loader = self.DataLoader
481
482 calib_dataloader_input_ids = data_loader(batch_encoded_input_ids, batch_size=batch_size, shuffle=False)
483 calib_dataloader_attention_mask = data_loader(
484 batch_encoded_attention_mask, batch_size=batch_size, shuffle=False
485 )
486
487 assert len(calib_dataloader_input_ids.dataset) == len(calib_dataloader_attention_mask.dataset)
488 assert len(calib_dataloader_input_ids) == len(calib_dataloader_attention_mask)
489
490 number_of_batched_samples = calib_size // batch_size
491
492 batched_input_ids = []
493 for idx, data in enumerate(calib_dataloader_input_ids):
494 batched_input_ids.append(data)
495 if idx == (number_of_batched_samples - 1):
496 break
497
498 batched_attention_mask = []
499 for idx, data in enumerate(calib_dataloader_attention_mask):
500 batched_attention_mask.append(data)
501 if idx == (number_of_batched_samples - 1):
502 break
503
504 print(
505 f"\n--Quantize-Script-- number_of_batched_samples={number_of_batched_samples}, "
506 f"batch-input-ids-list-len={len(batched_input_ids)}, batched_attention_mask={len(batched_attention_mask)}\n"
507 )
508
509 batched_inputs_list = []
510 for i in range(number_of_batched_samples):
511 input_ids = batched_input_ids[i]
512 attention_mask = batched_attention_mask[i]
513
514 inputs = self.make_model_input(
515 config,
516 input_ids,
517 attention_mask,
518 add_past_kv_inputs,
519 device,
520 use_fp16,
521 use_buffer_share,
522 add_position_ids,
523 )
524 inputs = {input_name: torch_tensor.cpu().numpy() for input_name, torch_tensor in inputs.items()}
525 batched_inputs_list.append(inputs)
526
527 print(f"\n--Quantize-Script-- number of batched inputs = {len(batched_inputs_list)}\n")
528 return batched_inputs_list
529
530
531def is_divisible(val1, val2):
532 return int(val2 * np.ceil(val1 / val2)) == val1
533
534
535class HQQWeightOnlyQuantizer:
536 def __init__(
537 self,
538 config: HQQWeightOnlyQuantConfig,
539 ):
540 self.config = config
541
542 # Proximal solver || weight - dequantize(quantize(weight))||_p^p
543 @staticmethod
544 def optimize_weights(
545 tensor,
546 scale,
547 zero,
548 min_max: list[int],
549 axis: int = 0,
550 opt_params: dict | None = None,
551 verbose=False,
552 ):
553 import torch # noqa: PLC0415
554
555 opt_params = {"lp_norm": 0.7, "beta": 1e1, "kappa": 1.01, "iters": 20} if opt_params is None else opt_params
556 lp_norm, beta, kappa, iters = (
557 opt_params["lp_norm"],
558 opt_params["beta"],
559 opt_params["kappa"],
560 opt_params["iters"],
561 )
562
563 dtype = torch.float16 if tensor.is_cuda else torch.float32
564 w_f = tensor.to(dtype)
565 scale = scale.to(dtype)
566 zero = zero.to(dtype)
567
568 def shrink_op(x, beta, p=lp_norm):
569 if p == 1:
570 return torch.sign(x) * torch.nn.functional.relu(torch.abs(x) - 1.0 / beta)
571 else:
572 return torch.sign(x) * torch.nn.functional.relu(
573 torch.abs(x) - (1.0 / beta) * torch.pow(torch.abs(x) + 1e-8, p - 1)
574 )
575
576 best_error = 1e4
577 for i in range(iters):
578 w_q = torch.round(w_f * scale + zero).clamp(min_max[0], min_max[1])
579 w_r = (w_q - zero) / scale
580 w_e = shrink_op(w_f - w_r, beta)
581 zero = torch.mean(w_q - (w_f - w_e) * scale, axis=axis, keepdim=True)
582 beta *= kappa
583
584 current_error = float(torch.abs(w_f - w_r).mean())
585 if verbose:
586 print(i, np.round(current_error, 6))
587 if current_error < best_error:
588 best_error = current_error
589 else:
590 break
591
592 del w_f, w_q, w_r, w_e
593
594 return scale, zero
595
596 @staticmethod
597 def pack_on_row_fast_248bit(pack_tensor, ori_int_tensor, bits):
598 if pack_tensor.shape[0] == ori_int_tensor.shape[0]:
599 ori_int_tensor = ori_int_tensor.T
600 pack_tensor = pack_tensor.T
601 if bits in [2, 4, 8]:
602 compress_ratio = pack_tensor.element_size() * 8 // bits
603 for j in range(compress_ratio):
604 pack_tensor[0:] |= ori_int_tensor[j::compress_ratio] << (bits * (j))
605 else:
606 raise NotImplementedError("Only 2,4,8 bits are supported.")
607
608 # from Official implementation of Half-Quadratic Quantization (HQQ)
609 def quantize_internal(
610 self, tensor, bits=4, channel_wise=True, group_size=64, optimize=True, round_zero=True, axis=1
611 ):
612 import torch # noqa: PLC0415
613
614 weight = tensor.float()
615 ori_shape = weight.shape
616
617 pad_len = (group_size - ori_shape[axis] % group_size) % group_size
618 if axis == 1:
619 weight = torch.nn.functional.pad(weight, (0, pad_len), "constant", 0)
620 else:
621 weight = torch.nn.functional.pad(weight, (0, 0, 0, pad_len), "constant", 0)
622 shape = weight.shape
623
624 # Reshape for grouping
625 if (group_size is not None) and channel_wise:
626 weight = weight.reshape([-1, group_size]) if (axis == 1) else weight.reshape([group_size, -1])
627
628 # Get min/max values
629 if channel_wise is False:
630 _min, _max = weight.min(), weight.max()
631 optimize = False
632 else:
633 _min = weight.min(axis=axis, keepdim=True)[0]
634 _max = weight.max(axis=axis, keepdim=True)[0]
635
636 max_v = 2**bits - 1
637 min_v = 0
638 min_max = [min_v, max_v]
639
640 # Note: here we work with the inverse of the scale to avoid division and quantize instead via weight*scale + zero, the scale is inverted later on.
641 # clamp to avoid half-precision problems
642 scale = (max_v / (_max - _min)).clamp(max=2e4)
643 #!!!!!!!!!!!!!!!
644 min_max_axis = _max - _min
645 if (min_max_axis == 0).sum().item() > 0:
646 min_max_axis[min_max_axis == 0] = max_v
647 scale = (max_v / min_max_axis).clamp(max=2e4)
648 zero = -_min * scale
649
650 if round_zero:
651 zero = torch.round(zero)
652
653 # Fine-tune weights
654 if optimize:
655 scale, zero = self.optimize_weights(tensor=weight, scale=scale, zero=zero, min_max=min_max, axis=axis)
656
657 # Quantize
658 # Necessary for fake quantization backprop
659 w_q = torch.round(weight * scale + zero).clamp(min_max[0], min_max[1])
660 w_q = w_q.reshape(shape).int()
661
662 scale = 1.0 / scale
663 if axis == 1:
664 scale = scale.reshape(shape[0], -1)
665 zero = zero.reshape(shape[0], -1)
666 else:
667 scale = scale.reshape(-1, shape[-1])
668 zero = zero.reshape(-1, shape[-1])
669 # cleanup
670 del weight, _min, _max
671
672 return w_q, scale.to(tensor.dtype), zero.to(tensor.dtype)
673
674 def quantize(self, node: NodeProto, graph_stack: list[GraphProto]) -> list[NodeProto]:
675 """
676 Target node: QOperator node: QDQ nodes:
677 MatMul MatMulNBits DeQuantizeLinear -> MatMul
678 Gather GatherBlockQuantized Gather, Gather, Gather (optional) -> DequantizeLinear
679 If the node is target node with fp32 or fp16 const weight, quantize the weight to int4 and
680 return the new nodes.
681 If QOperator format, return the corresponding QOperator nodes.
682 If QDQ format, return the corresdponging QDQ nodes.
683 Gather (quantized data) + Gather (scales) + Gather (optional, zero points) -> DequantizeLinear is
684 not supported yet because Gather does not support int4 data.
685 """
686 # With HQQ, zero points are in float. Current GatherBlockQuantized does not support float zero points.
687 if node.op_type == "Gather":
688 raise NotImplementedError("Gather quantization is not supported yet in HQQ")
689
690 import torch # noqa: PLC0415
691
692 logger.info(f"start to quantize {node.name} ...")
693 input_b = node.input[1]
694 b_pb, bs_graph = get_initializer(input_b, graph_stack)
695 if b_pb is None:
696 logger.info("MatMul doesn't have const weight. Skip to quantize")
697 return [node] # only care about constant weight
698
699 b_array = onnx.numpy_helper.to_array(b_pb)
700 if len(b_array.shape) != 2:
701 logger.info("MatMul weight is not 2D. Skip to quantize")
702 return [node] # can only process 2-D matrix
703 b_array_torch = torch.from_numpy(b_array)
704 if torch.cuda.is_available():
705 b_array_torch = b_array_torch.cuda()
706
707 bits = self.config.bits
708 quant_weight_torch, scales_torch, zero_points_torch = self.quantize_internal(
709 b_array_torch.T, bits=bits, group_size=self.config.block_size
710 )
711 quant_weight_torch = quant_weight_torch.contiguous()
712 scales_torch = scales_torch.contiguous()
713 zero_points_torch = zero_points_torch.contiguous()
714
715 packed_size = 8 // bits # number of elements packed into one byte
716
717 packed_torch = torch.zeros(
718 (quant_weight_torch.shape[0], quant_weight_torch.shape[1] // packed_size),
719 dtype=torch.uint8,
720 device=quant_weight_torch.device,
721 )
722 self.pack_on_row_fast_248bit(packed_torch, quant_weight_torch, bits)
723 scales = scales_torch.cpu().numpy()
724 zero_points = zero_points_torch.cpu().numpy()
725 # reshape to the predefined shape in MatmulNbits
726 scales = scales.reshape(-1)
727 zero_points = zero_points.reshape(-1)
728 rows, cols = b_array_torch.shape
729 block_size = self.config.block_size
730 blob_size = block_size // packed_size
731 k_blocks = (rows + block_size - 1) // block_size
732 packed_torch = packed_torch.reshape(cols, k_blocks, blob_size)
733
734 b_quant = onnx.numpy_helper.from_array(packed_torch.cpu().numpy())
735 b_quant.name = b_pb.name + "_Q" + str(bits)
736 for input in bs_graph.input:
737 if input.name == input_b:
738 bs_graph.input.remove(input)
739 break
740
741 scales_tensor = onnx.numpy_helper.from_array(scales)
742 scales_tensor.name = b_pb.name + "_scales"
743 bs_graph.initializer.extend([b_quant, scales_tensor])
744
745 input_names = [node.input[0], b_quant.name, scales_tensor.name]
746 zp_tensor = onnx.numpy_helper.from_array(zero_points)
747 zp_tensor.name = b_pb.name + "_zero_points"
748 bs_graph.initializer.extend([zp_tensor])
749 input_names.append(zp_tensor.name)
750
751 kwargs = {}
752 rows, cols = b_array.shape
753 kwargs["K"] = rows
754 kwargs["N"] = cols
755 kwargs["bits"] = bits
756 kwargs["block_size"] = self.config.block_size
757
758 matmul_q_node = onnx.helper.make_node(
759 "MatMulNBits",
760 inputs=input_names,
761 outputs=[node.output[0]],
762 name=node.name + "_Q" + str(bits) if node.name else "",
763 domain="com.microsoft",
764 **kwargs,
765 )
766
767 logger.info(f"complete quantization of {node.name} ...")
768
769 return [matmul_q_node]
770
771
772def get_initializer(name, graph_path: list[GraphProto]) -> tuple[TensorProto, GraphProto]:
773 for gid in range(len(graph_path) - 1, -1, -1):
774 graph = graph_path[gid]
775 for tensor in graph.initializer:
776 if tensor.name == name:
777 return tensor, graph
778 return None, None
779
780
781# transpose int4 matrix (packed as uint8)
782def transpose_packed_int4_matrix(packed, rows, cols):
783 # unpack to int4 matrix
784 total = rows * cols
785 high = (packed >> 4) & 0x0F
786 low = packed & 0x0F
787 int4_vals = np.empty(total, dtype=np.uint8)
788 int4_vals[0::2] = low
789 int4_vals[1::2] = high
790 int4_matrix = int4_vals.reshape((rows, cols))
791
792 # transpose int4 matrix
793 int4_matrix_transposed = int4_matrix.T
794
795 # pack to uint8
796 flat = int4_matrix_transposed.reshape(-1)
797 packed = ((flat[1::2] << 4) & 0xF0) | (flat[0::2] & 0x0F)
798 return packed.astype(np.uint8)
799
800
801class DefaultWeightOnlyQuantizer:
802 def __init__(self, config: DefaultWeightOnlyQuantConfig):
803 self.config = config
804
805 def qbits_block_quant(self, fp32weight: npt.ArrayLike) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
806 """4b/8b quantize fp32 weight to int4 using C++ kernels."""
807
808 qbits = self.config.bits
809 kpack = 8 // qbits
810 if len(fp32weight.shape) != 2:
811 raise ValueError("Current int4 block quantization only supports 2D tensors!")
812 rows, cols = fp32weight.shape
813
814 block_size = self.config.block_size
815 k_blocks = (rows + block_size - 1) // block_size
816
817 if self.config.quant_format == QuantFormat.QOperator:
818 blob_size = (block_size + kpack - 1) // kpack
819 padded_rows = k_blocks * block_size
820 pad_len = padded_rows - rows
821 if pad_len > 0:
822 fp32weight = np.pad(fp32weight, ((0, pad_len), (0, 0)), "constant")
823
824 # block wise quantization, each block comes from a single column
825 packed = np.zeros((cols, k_blocks, blob_size), dtype="uint8")
826 zero_point = np.zeros((cols, ((k_blocks + kpack - 1) // kpack)), dtype="uint8")
827 scales = np.zeros((cols, k_blocks), dtype=fp32weight.dtype)
828 if qbits == 2:
829 quantize_matmul_2bits(
830 packed, fp32weight, scales, zero_point, block_size, cols, rows, self.config.is_symmetric
831 )
832 elif qbits == 8:
833 quantize_matmul_8bits(
834 packed, fp32weight, scales, zero_point, block_size, cols, rows, self.config.is_symmetric
835 )
836 else:
837 quantize_matmul_4bits(
838 packed, fp32weight, scales, zero_point, block_size, cols, rows, self.config.is_symmetric
839 )
840 else:
841 # block size equal to rows (K) if channel wised quantize enabled
842 block_size = rows if self.config.channel_wised_quantize else self.config.block_size
843 k_blocks = (rows + block_size - 1) // block_size
844
845 assert qbits == 4, "QDQ format only support 4 bits quantization"
846 packed = np.zeros((rows * cols + 1) // 2, dtype="uint8")
847 zero_point = np.zeros((cols * k_blocks + 1) // 2, dtype="uint8")
848 scales = np.zeros((k_blocks, cols), dtype=fp32weight.dtype)
849 quantize_qdq_matmul_4bits(
850 packed, fp32weight, scales, zero_point, block_size, cols, rows, self.config.is_symmetric
851 )
852
853 return (packed, scales, zero_point)
854
855 def quantize_matmul(self, node: NodeProto, graph_stack: list[GraphProto]) -> list[NodeProto]:
856 """
857 Quantize weight B of MatMul node to int4 or int8.
858 Currently only support 2D constant matrix and axis 0 blockwise quantization.
859 """
860 bits = self.config.bits
861 if bits == 8:
862 qtype = TensorProto.INT8 if self.config.is_symmetric else TensorProto.UINT8
863 else:
864 qtype = TensorProto.INT4 if self.config.is_symmetric else TensorProto.UINT4
865 input_b = node.input[1]
866 b_tensor, b_graph = get_initializer(input_b, graph_stack)
867 if b_tensor is None:
868 logger.info("MatMul doesn't have const weight. Skip to quantize")
869 return [node] # only care about constant weight
870
871 b_ndarray = ir.from_proto(b_tensor).numpy()
872 if len(b_ndarray.shape) != 2:
873 logger.info("MatMul weight is not 2D. Skip to quantize")
874 return [node] # can only process 2-D matrix
875
876 bfloat16 = b_ndarray.dtype == "bfloat16"
877 if bfloat16:
878 b_ndarray = b_ndarray.astype(np.float32)
879
880 packed, scales, zero_points = self.qbits_block_quant(b_ndarray)
881 if bfloat16:
882 scales = scales.astype(ml_dtypes.bfloat16)
883
884 if self.config.quant_format == QuantFormat.QOperator:
885 b_quant = ir.serde.serialize_tensor(ir.Tensor(packed, name=b_tensor.name + f"_Q{bits}"))
886 scales_tensor = ir.serde.serialize_tensor(ir.Tensor(scales, name=b_tensor.name + "_scales"))
887 else:
888 b_quant = onnx.helper.make_tensor(
889 b_tensor.name + f"_DQ_Q{bits}", qtype, b_ndarray.shape, packed.tobytes(), True
890 )
891 scales_tensor = ir.serde.serialize_tensor(ir.Tensor(scales, name=b_tensor.name + "_DQ_scales"))
892
893 # if QDQ, CW and SYM enabled, optimize for Intel NPU, tranpose the weight to NHWC format will increase performance
894 qdq_opt_for_intel_npu_enabled = (
895 self.config.quant_format == QuantFormat.QDQ
896 and self.config.channel_wised_quantize
897 and self.config.is_symmetric
898 )
899 if qdq_opt_for_intel_npu_enabled:
900 rows, cols = b_ndarray.shape
901 packed = transpose_packed_int4_matrix(packed, rows, cols)
902 scales = scales.reshape((cols, 1)) # (cols, 1)
903 b_quant = onnx.helper.make_tensor(
904 b_tensor.name + f"_DQ_Q{bits}", qtype, [cols, rows], packed.tobytes(), True
905 )
906 scales_tensor = ir.serde.serialize_tensor(ir.Tensor(scales, name=b_tensor.name + "_DQ_scales"))
907
908 for input in b_graph.input:
909 if input.name == input_b:
910 b_graph.input.remove(input)
911 break
912
913 b_graph.initializer.extend([b_quant, scales_tensor])
914
915 output_nodes = []
916
917 if self.config.quant_format == QuantFormat.QOperator:
918 input_names = [node.input[0], b_quant.name, scales_tensor.name]
919 if not self.config.is_symmetric:
920 zp_tensor = onnx.numpy_helper.from_array(zero_points, b_tensor.name + "_zero_points")
921 input_names.append(zp_tensor.name)
922 b_graph.initializer.extend([zp_tensor])
923 kwargs = {}
924 rows, cols = b_ndarray.shape
925 kwargs["K"] = rows
926 kwargs["N"] = cols
927 kwargs["bits"] = bits
928 kwargs["block_size"] = self.config.block_size
929
930 # Do not output accuracy_level if it is 0 since the attribute is optional and is not supported by most EPs.
931 if self.config.accuracy_level:
932 kwargs["accuracy_level"] = self.config.accuracy_level
933
934 matmul_qbit_node = onnx.helper.make_node(
935 "MatMulNBits",
936 inputs=input_names,
937 outputs=[node.output[0]],
938 name=node.name + f"_Q{bits}" if node.name else "",
939 domain="com.microsoft",
940 **kwargs,
941 )
942
943 output_nodes.append(matmul_qbit_node)
944 else:
945 dq_input_names = [b_quant.name, scales_tensor.name]
946 dq_output_names = [b_quant.name + "_output"]
947 tp_input_names = [dq_output_names[0]]
948 tp_output_names = [dq_output_names[0] + "_transposed"]
949 matmul_input_names = [
950 node.input[0],
951 tp_output_names[0] if qdq_opt_for_intel_npu_enabled else dq_output_names[0],
952 ]
953 matmul_output_names = [node.output[0]]
954 if not self.config.is_symmetric:
955 zp_tensor = onnx.helper.make_tensor(
956 b_tensor.name + "_DQ_zero_points", qtype, scales.shape, zero_points.tobytes(), True
957 )
958 dq_input_names.append(zp_tensor.name)
959 b_graph.initializer.extend([zp_tensor])
960 rows, cols = b_ndarray.shape
961 dq_kwargs = {
962 "axis": 1 if qdq_opt_for_intel_npu_enabled else 0,
963 "block_size": rows if self.config.channel_wised_quantize else self.config.block_size,
964 }
965 dq_node = onnx.helper.make_node(
966 "DequantizeLinear",
967 inputs=dq_input_names,
968 outputs=dq_output_names,
969 name=node.name + f"_DQ_Q{bits}" if node.name else "",
970 **dq_kwargs,
971 )
972 matmul_node = onnx.helper.make_node(
973 "MatMul",
974 inputs=matmul_input_names,
975 outputs=matmul_output_names,
976 name=node.name + f"_matmul_Q{bits}" if node.name else "",
977 )
978 if qdq_opt_for_intel_npu_enabled:
979 tp_node = onnx.helper.make_node(
980 "Transpose",
981 inputs=tp_input_names,
982 outputs=tp_output_names,
983 perm=[1, 0],
984 )
985 output_nodes.extend([dq_node, tp_node, matmul_node])
986 else:
987 output_nodes.extend([dq_node, matmul_node])
988
989 return output_nodes
990
991 @staticmethod
992 def quant_slice_symmetric(data: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
993 max_val = np.max(data, axis=1, keepdims=True)
994 min_val = np.min(data, axis=1, keepdims=True)
995 abs_max = np.where(np.abs(max_val) > np.abs(min_val), max_val, min_val)
996
997 scale = abs_max / -8.0 # if max == min, max may be clipped
998 quantized_slice = np.where(scale == 0, 0, data / scale).round().clip(-8, 7).astype(np.int8)
999
1000 return quantized_slice, scale
1001
1002 @staticmethod
1003 def quant_slice_asymmetric(data: np.ndarray) -> tuple[np.ndarray, np.ndarray, np.ndarray]:
1004 min_val = np.minimum(data.min(axis=1, keepdims=True), 0)
1005 max_val = np.maximum(data.max(axis=1, keepdims=True), 0)
1006
1007 scale = (max_val - min_val) / 15.0
1008 zero_point = np.where(scale == 0, 8, -min_val / scale).round().clip(0, 15).astype(np.uint8)
1009 quantized_slice = np.where(scale == 0, 8, data / scale + zero_point).round().clip(0, 15).astype(np.uint8)
1010
1011 return quantized_slice, scale, zero_point
1012
1013 @staticmethod
1014 def pack_int8_to_int4(data: np.ndarray) -> np.ndarray:
1015 """Pack int8 data to int4 and store in uint8 ndarray."""
1016 data_flat = data.reshape(-1)
1017 if len(data_flat) % 2 != 0:
1018 data_flat = np.append(data_flat, 0)
1019 quant_data_int4 = (data_flat[::2] & 0xF) | ((data_flat[1::2] & 0xF) << 4)
1020
1021 return quant_data_int4.astype("uint8")
1022
1023 @staticmethod
1024 def quantize_ndarray(
1025 data: np.ndarray,
1026 quantize_axis: int,
1027 block_size: int,
1028 is_symmetric: bool,
1029 ) -> tuple[np.ndarray, np.ndarray, np.ndarray | None]:
1030 """Quantize ndarray data to int4 using numpy, return (quantized data, scales, zero points)."""
1031 # Get the shape of the matrix
1032 m = 1 # dimension of the matrix before the quantize axis
1033 k = data.shape[quantize_axis] # dimension of the matrix along the quantize axis
1034 n = 1 # dimension of the matrix after the quantize axis
1035 for i, dim in enumerate(data.shape):
1036 if i < quantize_axis:
1037 m *= dim
1038 elif i > quantize_axis:
1039 n *= dim
1040
1041 k_blocks = (k + block_size - 1) // block_size
1042 scales_shape = list(data.shape)
1043 scales_shape[quantize_axis] = k_blocks
1044
1045 data_reshape = data.reshape((m, k, n))
1046 scales = np.zeros((m, k_blocks, n), dtype=data.dtype)
1047 if is_symmetric:
1048 quant_data_int8 = np.zeros((m, k, n), dtype="int8")
1049 else:
1050 quant_data_int8 = np.zeros((m, k, n), dtype="uint8")
1051 zero_point_int8 = np.zeros((m, k_blocks, n), dtype="uint8")
1052
1053 # slice and quantize
1054 for i in range(0, k, block_size):
1055 end_idx = min(i + block_size, k)
1056 slice = data_reshape[:, i:end_idx, :]
1057
1058 if is_symmetric:
1059 quantized_slice_int8, scale_slice = DefaultWeightOnlyQuantizer.quant_slice_symmetric(slice)
1060 else:
1061 quantized_slice_int8, scale_slice, zero_point_slice_int8 = (
1062 DefaultWeightOnlyQuantizer.quant_slice_asymmetric(slice)
1063 )
1064
1065 quant_data_int8[:, i:end_idx, :] = quantized_slice_int8
1066 j = i // block_size
1067 scales[:, j : (j + 1), :] = scale_slice
1068 if not is_symmetric:
1069 zero_point_int8[:, j : (j + 1), :] = zero_point_slice_int8
1070
1071 # pack int8 to int4
1072 quant_data_int4 = DefaultWeightOnlyQuantizer.pack_int8_to_int4(quant_data_int8)
1073 zero_point_int4 = None
1074 if not is_symmetric:
1075 zero_point_int4 = DefaultWeightOnlyQuantizer.pack_int8_to_int4(zero_point_int8)
1076 scales = scales.reshape(scales_shape)
1077 return quant_data_int4, scales, zero_point_int4
1078
1079 def quantize_gather(self, node: NodeProto, graph_stack: list[GraphProto]) -> list[NodeProto]:
1080 """Quantize weight data of Gather node to int4."""
1081 assert self.config.quant_format == QuantFormat.QOperator, "Gather only supports QOperator format currently."
1082
1083 qtype = TensorProto.INT4 if self.config.is_symmetric else TensorProto.UINT4
1084 data_arg = node.input[0]
1085 data_tensorproto, data_graphproto = get_initializer(data_arg, graph_stack)
1086 if data_tensorproto is None:
1087 logger.info("Gather doesn't have const weight. Skip quantization.")
1088 return [node] # only care about constant weight
1089
1090 data_ndarray = onnx.numpy_helper.to_array(data_tensorproto)
1091 data_rank = len(data_ndarray.shape)
1092 quantize_axis = self.config.quant_axes.get("Gather", 1)
1093 block_size = self.config.block_size
1094
1095 assert quantize_axis < data_rank and quantize_axis >= -data_rank, "Invalid quantize axis for Gather node."
1096 assert block_size >= 16 and ((block_size - 1) & block_size == 0), "Invalid block size for Gather node."
1097
1098 quantize_axis = (quantize_axis + data_rank) % data_rank
1099 quantized_data, scales, zero_points = self.quantize_ndarray(
1100 data_ndarray, quantize_axis, block_size, self.config.is_symmetric
1101 )
1102
1103 for input in data_graphproto.input:
1104 if input.name == data_arg:
1105 data_graphproto.input.remove(input)
1106 break
1107
1108 quantized_data_tensorproto = onnx.helper.make_tensor(
1109 data_tensorproto.name + "_Q4", qtype, data_ndarray.shape, quantized_data.tobytes(), True
1110 )
1111 scales_tensorproto = onnx.numpy_helper.from_array(scales, data_tensorproto.name + "_scales")
1112 input_names = [quantized_data_tensorproto.name, node.input[1], scales_tensorproto.name]
1113 data_graphproto.initializer.extend([quantized_data_tensorproto, scales_tensorproto])
1114 if not self.config.is_symmetric:
1115 zp_tensorproto = onnx.helper.make_tensor(
1116 data_tensorproto.name + "_zero_points", qtype, scales.shape, zero_points.tobytes(), True
1117 )
1118 input_names.append(zp_tensorproto.name)
1119 data_graphproto.initializer.extend([zp_tensorproto])
1120
1121 try:
1122 gather_axis = onnx.helper.get_node_attr_value(node, "axis")
1123 except ValueError:
1124 gather_axis = 0
1125
1126 kwargs = {
1127 "gather_axis": gather_axis,
1128 "quantize_axis": quantize_axis,
1129 "block_size": block_size,
1130 }
1131
1132 gather_q4_node = onnx.helper.make_node(
1133 "GatherBlockQuantized",
1134 inputs=input_names,
1135 outputs=[node.output[0]],
1136 name=node.name + "_Q4" if node.name else "",
1137 domain="com.microsoft",
1138 **kwargs,
1139 )
1140
1141 return [gather_q4_node]
1142
1143 def quantize(self, node: NodeProto, graph_stack: list[GraphProto]) -> list[NodeProto]:
1144 """
1145 Target node: QOperator node: QDQ nodes:
1146 MatMul MatMulNBits DeQuantizeLinear -> MatMul
1147 Gather GatherBlockQuantized Gather, Gather, Gather (optional) -> DequantizeLinear
1148 If the node is target node with fp32 or fp16 const weight, quantize the weight to int4 and
1149 return the new nodes.
1150 If QOperator format, return the corresponding QOperator nodes.
1151 If QDQ format, return the corresdponging QDQ nodes.
1152 Gather (quantized data) + Gather (scales) + Gather (optional, zero points) -> DequantizeLinear is
1153 not supported yet because Gather does not support int4 data.
1154 """
1155 logger.info(f"start to quantize {node.name} ...")
1156
1157 bits = self.config.bits
1158 if node.op_type == "MatMul":
1159 if bits == 8 and self.config.quant_format == QuantFormat.QDQ:
1160 logger.error("MatMul only supports QOperator format for 8 bits quantization.")
1161 return [node]
1162 results = self.quantize_matmul(node, graph_stack)
1163 elif node.op_type == "Gather":
1164 if self.config.bits != 4:
1165 logger.error("Gather only supports 4 bits quantization.")
1166 return [node]
1167
1168 results = self.quantize_gather(node, graph_stack)
1169 else:
1170 logger.error(f"Unsupported operator {node.op_type} for weight only quantization. Skip quantization.")
1171 return [node]
1172
1173 logger.info(f"complete quantization of {node.name} with {self.config.bits} bits ...")
1174 return results
1175
1176
1177class NVAWQWeightOnlyQuantizer:
1178 def __init__(
1179 self,
1180 config: NVAWQWeightOnlyQuantConfig,
1181 ):
1182 self.config = config
1183
1184 def quantize_awq(self, model: ModelProto | str) -> ModelProto:
1185 """
1186 Perform nvidia_awq quantization using ModelOpt's int4 quantize function.
1187
1188 Args:
1189 model (ModelProto): The ONNX model to quantize.
1190
1191 Returns:
1192 ModelProto: The quantized ONNX model.
1193 """
1194 try:
1195 from modelopt.onnx.quantization.int4 import quantize as quantize_int4 # noqa: PLC0415
1196 except ImportError:
1197 print(
1198 "Please ensure that the 'modelopt' package is installed. Please install it using pip install nvidia_modelopt."
1199 )
1200 raise ImportError(
