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))))

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!
When I visualize the computation graph
g = dfgraph_from_tf_function(fn)(of, say, the training graph of a 3-layer feed-forward network) intf2.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))I checked the scheduling result
sched_result.scheduleintf2.wrapper.compile_tf2(), and indeed it is possible for parts of the backward pass to be scheduled before the end of the forward pass.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!