class ParallelDense(Layer):
def __init__(self, units, **kwargs):
super().__init__(**kwargs)
self.units = units
def build(self, input_shape):
super().build(input_shape)
self.kernel = self.add_weight(shape=[input_shape[1], input_shape[2], self.units], trainnable=True, initializer='glorot_uniform')
self.bias = self.add_weight(shape=[input_shape[1], self.units], trainnable=True, initializer='glorot_uniform')
def call(self, inputs):
return tf.einsum("bml, mlk -> bmk", inputs, self.kernel) + self.bias
This : Efficiently use Dense layers in parallel