(
self, quant_params: QuantParams, xnn_graph: XNNGraph, external_tag: str = None
)
| 275 | return dtype |
| 276 | |
| 277 | def get_quant_params( |
| 278 | self, quant_params: QuantParams, xnn_graph: XNNGraph, external_tag: str = None |
| 279 | ) -> XNNQuantParams: |
| 280 | if quant_params.per_channel: |
| 281 | scale = cast(torch.Tensor, quant_params.scale) |
| 282 | buffer_idx = len(xnn_graph.constant_data) |
| 283 | num_scales = scale.numel() |
| 284 | |
| 285 | if quant_params.per_channel_group: |
| 286 | scale = scale.to(torch.bfloat16) |
| 287 | |
| 288 | num_bytes = scale.untyped_storage().nbytes() |
| 289 | scale_array = ctypes.cast( |
| 290 | scale.untyped_storage().data_ptr(), |
| 291 | ctypes.POINTER(ctypes.c_char * num_bytes), |
| 292 | ).contents |
| 293 | scale_name = hashlib.sha256(bytes(scale_array)).hexdigest() |
| 294 | scale_name = "scale_" + scale_name |
| 295 | xnn_graph.constant_data.append( |
| 296 | ConstantDataOffset( |
| 297 | offset=UINT64_MAX, size=num_bytes, named_key=scale_name |
| 298 | ) |
| 299 | ) |
| 300 | if external_tag is not None: |
| 301 | logging.info( |
| 302 | f"Adding constant data with name, key {scale_name} and external_tag {external_tag} to named_data_store" |
| 303 | ) |
| 304 | self._named_data_store.add_named_data( |
| 305 | scale_name, bytes(scale_array), CONSTANT_TENSOR_ALIGNMENT, external_tag |
| 306 | ) |
| 307 | |
| 308 | if quant_params.per_channel_group: |
| 309 | return PerChannelGroupQuant( |
| 310 | scale=[], |
| 311 | channel_dim=quant_params.axis, |
| 312 | group_size=quant_params.group_size, |
| 313 | scale_buffer_idx=buffer_idx, |
| 314 | num_scales=num_scales, |
| 315 | ) |
| 316 | else: |
| 317 | return PerChannelQuant( |
| 318 | scale=[], |
| 319 | channel_dim=quant_params.axis, |
| 320 | scale_buffer_idx=buffer_idx, |
| 321 | num_scales=num_scales, |
| 322 | ) |
| 323 | elif quant_params.is_dynamic: |
| 324 | # NB: |
| 325 | # We use per_token quantization for per_tensor quantization |
| 326 | # Beacuase that's the only option in XNNPACK in absance of per_tensor dynamic quantization |
| 327 | # TODO: Upstream support for per_tensor dynamic quantization or broadcasting same scale value internally |
| 328 | return PerTokenDynamicQuant( |
| 329 | num_nonbatch_dims=quant_params.num_nonbatch_dims, |
| 330 | ) |
| 331 | |
| 332 | return PerTensorQuant( |
| 333 | scale=cast(float, quant_params.scale), |
| 334 | zero_point=cast(int, quant_params.zp), |
no test coverage detected