Conditional branches using tf.function

Viewed 215

I am having problems calculating the gradient using the gradient tape when using a tf.function with conditional branches.

Inside a gradient tape scope, I am trying to calculate the gradient of z w.r.t self.LMn. This works perfectly fine when I do not annotate the function with @tf.function. The error originates in a subclassed tf.keras.layers.Layer call function:

def call(self, x, training=None, labels=None, cur_switch=None, basis_filters=None):
    if basis_filters is None:
        basis_filters = self.initial_basis_filters

    gap = tf.reduce_mean(x, axis=[1, 2])
    z = tf.einsum('bc,cd->bd', gap, self.LMn)

    # ... code for out

    return out, z

The actual error is given as follows:

....
C:\...\ops\cond_v2.py:387 <lambda>
    lambda: _grad_fn(func_graph, grads), [], {},
C:\...\ops\cond_v2.py:363 _grad_fn
    assert len(func_graph.outputs) == len(grads)

AssertionError: 

More specifically,

func_graphs.outputs = [<tf.Tensor 'vgg16/block1a/block1a_conv/cond_2/Identity:0' shape=(128, 32, 32, None) dtype=float32>, <tf.Tensor 'vgg16/block1a/block1a_conv/cond_2/OptionalFromValue:0' shape=() dtype=variant>, <tf.Tensor 'vgg16/block1a/block1a_conv/cond_2/OptionalFromValue_1:0' shape=() dtype=variant>]

and

grads = (<tf.Tensor 'gradient_tape/vgg16/block1a/block1a_conv/strided_slice_3/StridedSliceGrad_1:0' shape=(128, 32, 32, None) dtype=float32>,)

I can assume that each of these outputs corresponds to the cases for some set of conditional branches outputs that are then fed into the tf.einsum function. I have read over all of the edge cases and precautions in the gradient tape documentation as this seems to the problem. Just a note that I am only performing conditional computation using hyperparameters (passed into tf.function as pythonic variables, such as basis_filters). There are also some conditional branches using (tensorflow ops) functions of these hyperparameters, is this allowed? or do I need to compute these values outside tf.function and pass these in as pythonic variables too?

I know the question is not completely clear and I can provide any extra information if needed. It would very helpful to have some guidance on what to look for with this kind of problem.

Thanks!

0 Answers
Related