I'd like to have a Conv2d layer, that is able to return a gradient with values at certain indices multiplied with 0.2. The original functionality should remain.
My idea was to subclass the Conv2d class and have a setter method to save the indices where the groundtruth is 0:
import tensorflow as tf
class BackpassFilter(tf.keras.layers.Layer):
def __init__(self, kernel_initializer, num_filter=48, kernel_size=3, stride=1, padding='same', name='keypoints_3_4'):
super(BackpassFilter, self).__init__()
self.fc = tf.keras.layers.Conv2D(kernel_initializer, num_filter, kernel_size, stride, padding, name, kernel_initializer=kernel_initializer)
def call(self, input):
return self.fc(input)
def add_filter_indices(self, ground_truth):
self.tmp_indices = tf.where(tf.equal(ground_truth, 0))
@tf.custom_gradient
def custom_gradient(self, x):
...
I have found that @tf.custom_gradient is something I could use. How can I implement the method to access the original gradients, manipulate and return them? Do I have to register the custom gradient somehow to be used in the class?