(expr)
| 3236 | """Traverse the expression graph and apply fusion""" |
| 3237 | |
| 3238 | def _fusion_pass(expr): |
| 3239 | # Full pass to find global dependencies |
| 3240 | seen = set() |
| 3241 | stack = [expr] |
| 3242 | dependents = defaultdict(set) |
| 3243 | dependencies = {} |
| 3244 | expr_mapping = {} |
| 3245 | |
| 3246 | while stack: |
| 3247 | next = stack.pop() |
| 3248 | |
| 3249 | if next._name in seen: |
| 3250 | continue |
| 3251 | seen.add(next._name) |
| 3252 | |
| 3253 | if is_valid_blockwise_op(next): |
| 3254 | dependencies[next._name] = set() |
| 3255 | if next._name not in dependents: |
| 3256 | dependents[next._name] = set() |
| 3257 | expr_mapping[next._name] = next |
| 3258 | |
| 3259 | for operand in next.dependencies(): |
| 3260 | stack.append(operand) |
| 3261 | if is_valid_blockwise_op(operand): |
| 3262 | if next._name in dependencies: |
| 3263 | dependencies[next._name].add(operand._name) |
| 3264 | dependents[operand._name].add(next._name) |
| 3265 | expr_mapping[operand._name] = operand |
| 3266 | expr_mapping[next._name] = next |
| 3267 | |
| 3268 | # Traverse each "root" until we find a fusable sub-group. |
| 3269 | # Here we use root to refer to a Blockwise Expr node that |
| 3270 | # has no Blockwise dependents |
| 3271 | roots = [ |
| 3272 | expr_mapping[k] |
| 3273 | for k, v in dependents.items() |
| 3274 | if v == set() |
| 3275 | or all(not is_valid_blockwise_op(expr_mapping[_expr]) for _expr in v) |
| 3276 | ] |
| 3277 | while roots: |
| 3278 | root = roots.pop() |
| 3279 | seen = set() |
| 3280 | stack = [root] |
| 3281 | group = [] |
| 3282 | while stack: |
| 3283 | next = stack.pop() |
| 3284 | |
| 3285 | if next._name in seen: |
| 3286 | continue |
| 3287 | seen.add(next._name) |
| 3288 | |
| 3289 | group.append(next) |
| 3290 | for dep_name in dependencies[next._name]: |
| 3291 | dep = expr_mapping[dep_name] |
| 3292 | |
| 3293 | stack_names = {s._name for s in stack} |
| 3294 | group_names = {g._name for g in group} |
| 3295 | if ( |
no test coverage detected