import matplotlib.pyplot as plt # Instantiate a PositionalEncoder class d_model = 400 max_length = 100 rounding = 4 PE = PositionalEncoder(d_model=d_model, max_length=max_length, rounding=rounding) # Generate positional encodings input = np.round(np.random.rand(max_length, d_model), 4) positional_encoding = PE.generate_positional_encoding() # Plot positional encodings cax = plt.matshow(positional_encoding, cmap='coolwarm') plt.title(f'Positional Encoding Matrix ({d_model=}, {max_length=})') plt.ylabel('Position of the Embeddingnin the Sequence, pos') plt.xlabel('Embedding Dimension, i') plt.gcf().colorbar(cax) plt.gca().xaxis.set_ticks_position('bottom')