| 195 | return knn_probs |
| 196 | |
| 197 | def knns_to_log_prob(self, knns, neg_dists): |
| 198 | probs = torch.nn.functional.softmax(neg_dists / self.knn_temperature, dim=-1) |
| 199 | vals_at_knns = self.vals[knns].squeeze(-1) # (nonpad batch * time, k) |
| 200 | knn_log_probs = torch.full(size=(vals_at_knns.shape[:-1] + (self.vocab_size,)), fill_value=0.0).to(self.device) \ |
| 201 | .scatter_add(dim=-1, index=vals_at_knns, src=probs).log() # (nonpad_batch * time, vocab) |
| 202 | knn_log_probs = torch.nan_to_num(knn_log_probs, nan=None, neginf=-10000.0) |
| 203 | return knn_log_probs |
| 204 | |
| 205 | def register_hook(self, layer, func, pre=False): |
| 206 | handle = layer.register_forward_pre_hook(func) if pre else layer.register_forward_hook(func) |