MCPcopy Create free account
hub / github.com/pytorch/executorch / from_q_dq_node

Method from_q_dq_node

backends/xnnpack/operators/quant_params.py:235–322  ·  view source on GitHub ↗
(
        cls, quant_node: torch.fx.Node, ep: Optional[ExportedProgram] = None
    )

Source from the content-addressed store, hash-verified

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 )

Callers 8

define_nodeMethod · 0.80
define_nodeMethod · 0.80
from_weightsMethod · 0.80
from_inputsMethod · 0.80
from_outputsMethod · 0.80
define_nodeMethod · 0.80
define_nodeMethod · 0.80
define_nodeMethod · 0.80

Calls 10

check_or_raiseFunction · 0.90
is_quantFunction · 0.90
is_dequantFunction · 0.90
is_dynamic_qdqFunction · 0.90
is_per_channelFunction · 0.90
is_per_channel_groupFunction · 0.90
is_affine_qdqFunction · 0.90
keysMethod · 0.80

Tested by

no test coverage detected