Mark all ops reached from "from_ops". Args: from_ops: list of Operations. reached_ops: set of Operations. func_graphs: list of FuncGraphs. This method will traverse through these functions if they capture from_ops or any reachable ops.
(from_ops, reached_ops, func_graphs)
| 49 | |
| 50 | |
| 51 | def _MarkReachedOps(from_ops, reached_ops, func_graphs): |
| 52 | """Mark all ops reached from "from_ops". |
| 53 | |
| 54 | Args: |
| 55 | from_ops: list of Operations. |
| 56 | reached_ops: set of Operations. |
| 57 | func_graphs: list of FuncGraphs. This method will traverse through |
| 58 | these functions if they capture from_ops or any reachable ops. |
| 59 | """ |
| 60 | queue = collections.deque() |
| 61 | queue.extend(from_ops) |
| 62 | while queue: |
| 63 | op = queue.popleft() |
| 64 | if op not in reached_ops: |
| 65 | reached_ops.add(op) |
| 66 | for output in op.outputs: |
| 67 | if _IsBackpropagatable(output): |
| 68 | queue.extend(_Consumers(output, func_graphs)) |
| 69 | |
| 70 | |
| 71 | def _PendingCount(to_ops, from_ops, colocate_gradients_with_ops, func_graphs, |
no test coverage detected