logits: N x d
(logits)
| 81 | |
| 82 | |
| 83 | def softmax(logits): |
| 84 | """ |
| 85 | logits: N x d |
| 86 | """ |
| 87 | max_value = np.max(logits, axis=1, keepdims=True) |
| 88 | exp = np.exp(logits - max_value) |
| 89 | exp_sum = np.sum(exp, axis=1, keepdims=True) |
| 90 | dist = exp / exp_sum |
| 91 | return dist |
| 92 | |
| 93 | |
| 94 | def get_keep_pos_idxs(labels, remove_blank=None): |
no outgoing calls
no test coverage detected