Here is the code.
def add_temporal_block(previous, skip, dilation, cropping):
for _ in range(block_size):
convs = []
convs.append(Conv1D(fixed_filters, kernel_size, dilation_rate=(dilation,), padding='causal')(previous))
if len(convs) > 1:
previous = Concatenate(axis=-1)(convs)
else:
previous = convs[0]
previous = BatchNormalization(axis=2)(previous)
previous = PReLU(shared_axes=[1])(previous)
drop_left = block_size * (kernel_size - 1) * dilation
cropping += drop_left
if skip is None:
skip = Conv1D(fixed_filters, 1, padding='causal')(previous)
cropping_value = max(cropping - (receptive_field_size - 1), 0)
out = Add()([Cropping1D(cropping=(drop_left, 0))(previous), previous])
skip_out = Cropping1D(cropping=(cropping_value, 0))(out) if cropping_value > 0 else out
if skip is not None:
skip_out = Add()([skip, Conv1D(fixed_filters, 1, padding='causal')(skip_out)])
else:
skip_out = Conv1D(fixed_filters, 1, padding='causal')(skip_out)
return PReLU(shared_axes=[1])(out), skip_out, cropping
def TCN(input_dim):
dilations = [2 ** i for i in range(8)]
input_layer = Input(shape=(None, input_dim[1]))
cropping = 0
prev_layer, skip_layer, _ = add_temporal_block(input_layer, None, 1, cropping)
for dilation in dilations:
prev_layer, skip_layer, cropping = add_temporal_block(prev_layer, skip_layer, dilation, cropping)
output_layer = PReLU(shared_axes=[1])(skip_layer)
output_layer = Conv1D(fixed_filters, kernel_size=1, padding='causal')(output_layer)
output_layer = PReLU(shared_axes=[1])(output_layer)
output_layer = Conv1D(1, kernel_size=1, padding='causal')(output_layer)
return Model(input_layer, output_layer)