| 127 | |
| 128 | |
| 129 | def plot_lang_distribution(lang_distribution, candidate_langs, candidate_layers): |
| 130 | lang_distribution_matrix = [] |
| 131 | for layer in candidate_layers: |
| 132 | lang_distribution_matrix.append([lang_distribution[layer][lang] for lang in candidate_langs]) |
| 133 | lang_distribution_matrix = np.array(lang_distribution_matrix).T |
| 134 | fig, ax = plt.subplots(figsize=(11,3)) |
| 135 | cmap = sns.color_palette("ch:start=.2,rot=-.3", as_cmap=True) |
| 136 | sns.heatmap(lang_distribution_matrix, ax=ax, xticklabels=candidate_layers, yticklabels=candidate_langs, cmap=cmap) |
| 137 | plt.title('Layerwise Language Distribution') |
| 138 | plt.xlabel('Layer') |
| 139 | plt.ylabel('Language') |
| 140 | plt.show() |
| 141 | plt.savefig('lang_distribution.png') |
| 142 | |
| 143 | |
| 144 | |