Samples the top-k edges from a set of logits.
(scores, k, randomize=False)
| 217 | |
| 218 | |
| 219 | def 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 | |
| 232 | def sample_inplace_to_torch(sample): |
no test coverage detected