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

Method _parse_sdpa_args_and_kwargs

backends/mlx/patterns.py:415–422  ·  view source on GitHub ↗
(cls, sdpa_node: Node)

Source from the content-addressed store, hash-verified

413
414 @classmethod
415 def _parse_sdpa_args_and_kwargs(cls, sdpa_node: Node):
416 q, k, v = sdpa_node.args[0:3]
417 attn_mask = sdpa_node.args[3] if len(sdpa_node.args) > 3 else None
418 dropout_p = sdpa_node.args[4] if len(sdpa_node.args) > 4 else 0.0
419 is_causal = sdpa_node.args[5] if len(sdpa_node.args) > 5 else False
420 enable_gqa = sdpa_node.args[6] if len(sdpa_node.args) > 6 else False
421 scale = sdpa_node.kwargs.get("scale", None)
422 return q, k, v, attn_mask, dropout_p, is_causal, scale, enable_gqa
423
424 @classmethod
425 def _try_unwrap_repeat_kv(cls, node: Node) -> Optional[Tuple[Node, List[Node]]]:

Callers 2

maybe_createMethod · 0.80
__call__Method · 0.80

Calls 1

getMethod · 0.45

Tested by

no test coverage detected