(
cls, quant_node: torch.fx.Node, ep: Optional[ExportedProgram] = None
)
| 233 | |
| 234 | @classmethod |
| 235 | def from_q_dq_node( |
| 236 | cls, quant_node: torch.fx.Node, ep: Optional[ExportedProgram] = None |
| 237 | ) -> QuantParams: |
| 238 | check_or_raise( |
| 239 | is_quant(quant_node) or is_dequant(quant_node), |
| 240 | f"building quantizer from q/dq node but was given node:{quant_node}", |
| 241 | ) |
| 242 | q_input = quant_node.all_input_nodes[0] |
| 243 | |
| 244 | # TODO: Use presence of choose_qparam node to determine if this is a dynamic quantization |
| 245 | if is_dynamic_qdq(quant_node): |
| 246 | return cls._from_dynamic_input_node(quant_node) |
| 247 | |
| 248 | per_channel = is_per_channel(quant_node) |
| 249 | |
| 250 | _groupwise = is_per_channel_group(quant_node) |
| 251 | quant_node_args = quant_node.args |
| 252 | if _groupwise and is_affine_qdq(quant_node): |
| 253 | quant_node_args = extract_qdq_affine_op_args_for_decomposed_ops(quant_node) |
| 254 | |
| 255 | scale = quant_node_args[1] |
| 256 | zp = quant_node_args[2] |
| 257 | axis = 0 |
| 258 | if per_channel: |
| 259 | assert isinstance(scale, torch.fx.Node) and isinstance(scale.target, str) |
| 260 | assert isinstance(zp, torch.fx.Node) and isinstance(zp.target, str) |
| 261 | assert ( |
| 262 | ep is not None |
| 263 | ), "ExportedProgram must be provided to extract per channel params" |
| 264 | |
| 265 | def _get_tensor(node): |
| 266 | param = get_param_tensor(ep, node) |
| 267 | assert param is not None, f"Expected to find param tensor for {node}" |
| 268 | return cast(torch.Tensor, param) |
| 269 | |
| 270 | scale = _get_tensor(scale) |
| 271 | zp = _get_tensor(zp) |
| 272 | axis = cast(int, quant_node_args[3]) |
| 273 | |
| 274 | if _groupwise: |
| 275 | scale_tensor = cast(torch.Tensor, scale) |
| 276 | if scale_tensor.ndim == 1: |
| 277 | scale_tensor = scale_tensor.reshape(-1, 1) |
| 278 | zp = zp.reshape(-1, 1) |
| 279 | scale = scale_tensor |
| 280 | |
| 281 | assert ( |
| 282 | scale_tensor.ndim == 2 |
| 283 | ), "Weight scale must be 2D for per_channel_group [de]quant node, got {scale.ndim}D" |
| 284 | axis = 0 # axis is ignored for groupwise quantization |
| 285 | |
| 286 | check_or_raise( |
| 287 | bool( |
| 288 | quant_node_args[-1] != torch.uint8 |
| 289 | or quant_node_args[-1] != torch.quint8 |
| 290 | ), |
| 291 | "XNNPACK does not support unsigned quantization", |
| 292 | ) |
no test coverage detected