Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
matmul_nbits_quantizer.py1639 linesDownload Raw Back to quantization
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(

Showing the first 1,200 of 1639 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai