MCPcopy Create free account
hub / github.com/BioinfoMachineLearning/FlowDock / topk_edge_mask_from_logits

Function topk_edge_mask_from_logits

flowdock/utils/model_utils.py:219–229  ·  view source on GitHub ↗

Samples the top-k edges from a set of logits.

(scores, k, randomize=False)

Source from the content-addressed store, hash-verified

217
218
219def topk_edge_mask_from_logits(scores, k, randomize=False):
220 """Samples the top-k edges from a set of logits."""
221 assert len(scores.shape) == 3, "Scores should have shape [B, N, N]"
222 if randomize:
223 noise = torch.rand_like(scores)
224 scores = scores - torch.log(-torch.log(noise))
225 node_degree = min(k, scores.shape[2])
226 _, topk_idx = torch.topk(scores, node_degree, dim=-1, largest=True)
227 edge_mask = scores.new_zeros(scores.shape, dtype=torch.bool)
228 edge_mask = edge_mask.scatter_(dim=2, index=topk_idx, value=1).bool()
229 return edge_mask
230
231
232def sample_inplace_to_torch(sample):

Calls 1

logMethod · 0.80

Tested by

no test coverage detected