(activation, multimodality_times)
| 120 | |
| 121 | |
| 122 | def calculate_multimodality(activation, multimodality_times): |
| 123 | assert len(activation.shape) == 3 |
| 124 | assert activation.shape[1] > multimodality_times |
| 125 | num_per_sent = activation.shape[1] |
| 126 | |
| 127 | first_dices = np.random.choice(num_per_sent, multimodality_times, replace=False) |
| 128 | second_dices = np.random.choice(num_per_sent, multimodality_times, replace=False) |
| 129 | dist = linalg.norm(activation[:, first_dices] - activation[:, second_dices], axis=2) |
| 130 | return dist.mean() |