Skip to content

Invalid schedule due to incomplete computation graph #156

Description

@haoming-codes

When I visualize the computation graph g = dfgraph_from_tf_function(fn) (of, say, the training graph of a 3-layer feed-forward network) in tf2.wrapper.compile_tf2(), I notice that there isn't an edge between the loss computation and the backward pass. (See attached image; here, for simplicity of graphing, node 'Sum' is a dummy loss function, loss = tf.math.reduced_sum(predictions))

model = keras.Sequential([
    layers.Dense(10, name="L1", activation=None, use_bias=False),
    layers.Dense(10, name="L2", activation=None, use_bias=False),
    layers.Dense(10, name="L3", activation=None, use_bias=False),
    ])

@tf.function
def train_step(model, x):
   with tf.GradientTape() as tape:
       pred = model(x)
       loss = tf.math.reduce_sum(pred)
   gradients = tape.gradient(loss, model.trainable_variables)
   return pred, gradients

g = dfgraph_from_tf_function(train_step.get_concrete_function(model, tf.ones((5,10))))

image

I checked the scheduling result sched_result.schedule in tf2.wrapper.compile_tf2(), and indeed it is possible for parts of the backward pass to be scheduled before the end of the forward pass.

gradient_tape/Reshape
gradient_tape/Tile
sequential/L1/MatMul
gradient_tape/sequential/L3/MatMul/MatMul
sequential/L2/MatMul
gradient_tape/sequential/L2/MatMul/MatMul
sequential/L3/MatMul
gradient_tape/sequential/L1/MatMul/MatMul
Sum
gradient_tape/sequential/L2/MatMul/MatMul_1
gradient_tape/sequential/L3/MatMul/MatMul_1

I would like to know if this is by design, i.e. Tensorflow is expected to make adjustments to this (invalid) schedule during execution. Thank you!

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions