One possible implementation could be this :
class LocallyDenseLayer(tf.keras.layers.Layer):
def __init__(self, k, m, *args, **kwargs):
super().__init__(args, kwargs)
# alternatively, you can move that setup in the build method
# and infer the shape from the input
# this is left as an exercise to the reader
self.w = self.add_weight(name="weight", shape=(k,m))
def call(self, inputs):
# assuming input has shape [batch, k, m]
dotp = tf.linalg.diag_part(tf.tensordot(inputs, self.w, axes=[[1],[0]]))
return tf.nn.relu(dotp)
Using tf.tensordot to do the dot product over the dimension k and extracting only the diagonal, that contains what we want.
A simple example of usage :
X = tf.random.normal((100,5,1024))
y = tf.random.normal((100,1))
model = tf.keras.Sequential(
[
tf.keras.Input((5,1024)),
LocallyDenseLayer(5,1024),
tf.keras.layers.Dense(512, activation="relu"),
tf.keras.layers.Dense(1, activation="sigmoid")
]
)
model.compile(loss="mse",optimizer="sgd")
model.fit(X,y)