( # noqa: C901
simple_graph: SimpleGraph,
current_node: SimpleNode,
pattern_nodes: Dict[str, SimpleNode],
first_node: SimpleNode,
matched_name_map: Dict[str, str],
)
| 82 | |
| 83 | # [Pattern matching] |
| 84 | def check_inputs( # noqa: C901 |
| 85 | simple_graph: SimpleGraph, |
| 86 | current_node: SimpleNode, |
| 87 | pattern_nodes: Dict[str, SimpleNode], |
| 88 | first_node: SimpleNode, |
| 89 | matched_name_map: Dict[str, str], |
| 90 | ) -> bool: |
| 91 | # check op type |
| 92 | if first_node.op != "*": |
| 93 | matched_ops = [op.strip() for op in first_node.op.split("|")] |
| 94 | if current_node.op not in matched_ops: |
| 95 | return False |
| 96 | # check node name |
| 97 | if first_node.name in matched_name_map: |
| 98 | if matched_name_map[first_node.name] != current_node.name: |
| 99 | return False |
| 100 | # check inputs |
| 101 | if (len(first_node.inputs) == 1) and (first_node.inputs[0] == "*"): |
| 102 | matched_name_map[first_node.name] = current_node.name |
| 103 | return True |
| 104 | # if inputs contains both unknown inputs and known inputs |
| 105 | if (len(first_node.inputs) > 1) and ("*" in first_node.inputs): |
| 106 | known_inputs = [name for name in first_node.inputs if name != "*"] |
| 107 | start_idx = 0 |
| 108 | for key_name in known_inputs: |
| 109 | matched = False |
| 110 | if key_name.isdigit(): |
| 111 | matched = True |
| 112 | continue |
| 113 | for i in range(start_idx, len(current_node.inputs)): |
| 114 | input_name = current_node.inputs[i] |
| 115 | cur_input_node = simple_graph.get_simple_node_by_name(input_name) |
| 116 | expected_input_op_str = (pattern_nodes[key_name].op).strip() |
| 117 | if "|" in expected_input_op_str: |
| 118 | expected_input_ops = expected_input_op_str.split("|") |
| 119 | else: |
| 120 | expected_input_ops = list([expected_input_op_str]) |
| 121 | if (cur_input_node.op in expected_input_ops) and ( |
| 122 | check_inputs( |
| 123 | simple_graph, |
| 124 | cur_input_node, |
| 125 | pattern_nodes, |
| 126 | pattern_nodes[key_name], |
| 127 | matched_name_map, |
| 128 | ) |
| 129 | ): |
| 130 | matched = True |
| 131 | start_idx = i |
| 132 | if not matched: |
| 133 | return False |
| 134 | # if all listed inputs are known inputs |
| 135 | else: |
| 136 | if len(current_node.inputs) != len(first_node.inputs): |
| 137 | return False |
| 138 | for i, input_name in enumerate(current_node.inputs): |
| 139 | cur_input_node = simple_graph.get_simple_node_by_name(input_name) |
| 140 | if first_node.inputs[i].isdigit(): |
| 141 | continue |
no test coverage detected