import itertools # Iterate through all layers for idx_out, layer in enumerate(all_layers): # If layer is a convolutional filter if type(layer) == nn.Conv2d: # Print layer name print(f"n{idx_out+1}. Layer: {layer} n") # Prepare plot and weights plt.figure(figsize=(25,6)) weights = conv_weights[idx_out][:,0,:,:] # only first input channel weights = weights.detach().to('cpu') # Enumerate over filter weights (only first input channel) for idx_in, f in enumerate(weights): plt.subplot(2,8, idx_in+1) plt.imshow(f, cmap="gray") plt.title(f"Filter {idx_in+1}") # Print texts for i, j in itertools.product(range(f.shape[0]), range(f.shape[1])): if f[i,j] > f.mean(): color = 'black' else: color = 'white' plt.text(j, i, format(f[i, j], '.2f'), horizontalalignment='center', verticalalignment='center', color=color) plt.axis("off") plt.show() plt.close()