(cls, sdpa_node: Node)
| 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]]]: |
no test coverage detected