(grid_size=33, dim=512)
| 169 | |
| 170 | |
| 171 | def vis_position_embedding(grid_size=33, dim=512): |
| 172 | if False: |
| 173 | ref_x = grid_size // 2 |
| 174 | ref_y = grid_size // 2 |
| 175 | elif True: |
| 176 | ref_x = 0 |
| 177 | ref_y = 0 |
| 178 | pos_embed = get_2d_sincos_pos_embed(dim, grid_size) |
| 179 | pos_embed_3d = pos_embed.reshape(dim, grid_size, grid_size) |
| 180 | reference_pts = pos_embed_3d[:, ref_x : ref_x + 1, ref_y : ref_y + 1] |
| 181 | distance = np.linalg.norm(pos_embed_3d - reference_pts, ord=1, axis=0) |
| 182 | print(distance.shape) |
| 183 | plt.imshow(distance, cmap="inferno") |
| 184 | plt.colorbar() |
| 185 | plt.savefig("distance_pe_vis.png") |
| 186 | plt.show() |
| 187 | |
| 188 | |
| 189 | if __name__ == "__main__": |
no test coverage detected