I am trying to build a sparsely connected layer where each output neuron is connected to exactly n inputs.
I've come up with an implementation that seems to work. My model appears to converge when I replace the dense layer with the custom layer. I usually add significantly more neurons to compensate for the loss in predictive power due to limiting the number of connexions.
However, the implementation is excruciatingly slow (like 30x slower than with standard dense layers). Any idea where could I improve my code to make it 'work'? Or another approach that would lead to a sparsely connected layer?
class NConnected(tf.keras.layers.Layer):
#creator
def __init__(self, units=32, n_max = 5):
super(NConnected, self).__init__()
self.units = units
self.n_max = n_max
#creates weights
def build(self, input_shape):
self.w = self.add_weight(name='w',
shape=(self.n_max,self.units),
initializer='uniform',
trainable=True)
self.b = self.add_weight(name='b',
shape=self.units,
initializer='zeros',
trainable=True)
mask_t = []
for i in range(self.units):
a = range(1,input_shape[-1]+1)
r = random.sample(a,self.n_max)
mask_t.append(np.array([i in r for i in a]))
mask_t = tf.constant(np.array(mask_t))
self.mask = mask_t
self.built=True
#operation:
def call(self, inputs):
m = tf.map_fn(fn=lambda t: tf.boolean_mask(inputs,t,axis=1),elems = self.mask, fn_output_signature=tf.float32)
m = tf.transpose(m, [1, 0, 2])
res = tf.math.reduce_sum(tf.tensordot(m,self.w,axes=1),axis=2)+self.b
return res
#for saving the model - only necessary if you have parameters in __init__
def get_config(self):
config = super(NConnected, self).get_config()
return config
As mentionned the layer is way slower than an equivalent sized dense layer. A simple forward pass as an exemple :
A small batch of random data :
x_test = tf.random.normal((1000,100), mean=0.0, stddev=1.0, dtype=tf.dtypes.float32, seed=None, name=None)
Simple test loop :
%%time
test_n = NConnected(units=32, n_max = 3)
for i in range(100):
test_n(x_test)
returns :
CPU times: user 4.77 s, sys: 59.1 ms, total: 4.83 s Wall time: 4.72 s
While :
%%time
test_dense = layers.Dense(units=32)
for i in range(100):
test_dense(x_test)
Returns :
CPU times: user 47.9 ms, sys: 7.2 ms, total: 55.1 ms Wall time: 45.3 ms
So around 100x slower, which renders the use of that sparse layer near impossible in practice.