MCPcopy Create free account
hub / github.com/LUMIA-Group/MemoryDecoder / knns_to_log_prob

Method knns_to_log_prob

knn_utils/saveEmbedMulti.py:197–203  ·  view source on GitHub ↗
(self, knns, neg_dists)

Source from the content-addressed store, hash-verified

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)

Callers 1

post_forward_hookMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected