Do a forward graph walk and return all the visited ops. Args: seed_ops: an iterable of operations from which the forward graph walk starts. If a list of tensors is given instead, the seed_ops are set to be the consumers of those tensors. inclusive: if True the given seed_ops a
(seed_ops,
inclusive=True,
within_ops=None,
within_ops_fn=None,
stop_at_ts=(),
control_outputs=None)
| 385 | |
| 386 | |
| 387 | def get_forward_walk_ops(seed_ops, |
| 388 | inclusive=True, |
| 389 | within_ops=None, |
| 390 | within_ops_fn=None, |
| 391 | stop_at_ts=(), |
| 392 | control_outputs=None): |
| 393 | """Do a forward graph walk and return all the visited ops. |
| 394 | |
| 395 | Args: |
| 396 | seed_ops: an iterable of operations from which the forward graph |
| 397 | walk starts. If a list of tensors is given instead, the seed_ops are set |
| 398 | to be the consumers of those tensors. |
| 399 | inclusive: if True the given seed_ops are also part of the resulting set. |
| 400 | within_ops: an iterable of `tf.Operation` within which the search is |
| 401 | restricted. If `within_ops` is `None`, the search is performed within |
| 402 | the whole graph. |
| 403 | within_ops_fn: if provided, a function on ops that should return True iff |
| 404 | the op is within the graph traversal. This can be used along within_ops, |
| 405 | in which case an op is within if it is also in within_ops. |
| 406 | stop_at_ts: an iterable of tensors at which the graph walk stops. |
| 407 | control_outputs: a `util.ControlOutputs` instance or None. |
| 408 | If not `None`, it will be used while walking the graph forward. |
| 409 | Returns: |
| 410 | A Python set of all the `tf.Operation` ahead of `seed_ops`. |
| 411 | Raises: |
| 412 | TypeError: if `seed_ops` or `within_ops` cannot be converted to a list of |
| 413 | `tf.Operation`. |
| 414 | """ |
| 415 | _, control_outputs = check_cios(False, control_outputs) |
| 416 | if not util.is_iterable(seed_ops): |
| 417 | seed_ops = [seed_ops] |
| 418 | if not seed_ops: |
| 419 | return [] |
| 420 | if isinstance(seed_ops[0], tf_ops.Tensor): |
| 421 | ts = util.make_list_of_t(seed_ops, allow_graph=False) |
| 422 | seed_ops = util.get_consuming_ops(ts) |
| 423 | else: |
| 424 | seed_ops = util.make_list_of_op(seed_ops, allow_graph=False) |
| 425 | |
| 426 | seed_ops = frozenset(seed_ops) |
| 427 | stop_at_ts = frozenset(util.make_list_of_t(stop_at_ts)) |
| 428 | if within_ops: |
| 429 | within_ops = util.make_list_of_op(within_ops, allow_graph=False) |
| 430 | within_ops = frozenset(within_ops) |
| 431 | seed_ops &= within_ops |
| 432 | |
| 433 | def is_within(op): |
| 434 | return (within_ops is None or op in within_ops) and ( |
| 435 | within_ops_fn is None or within_ops_fn(op)) |
| 436 | |
| 437 | result = list(seed_ops) |
| 438 | wave = set(seed_ops) |
| 439 | while wave: |
| 440 | new_wave = set() |
| 441 | for op in wave: |
| 442 | for new_t in op.outputs: |
| 443 | if new_t in stop_at_ts: |
| 444 | continue |
no test coverage detected