MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / get_forward_walk_ops

Function get_forward_walk_ops

tensorflow/contrib/graph_editor/select.py:387–456  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

385
386
387def 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

Callers 2

get_walks_union_opsFunction · 0.85

Calls 5

check_ciosFunction · 0.85
is_withinFunction · 0.70
consumersMethod · 0.45
addMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected