Generate node IDs randomly for mixup; avoid mixup the same node. Args: x: The latent embedding or node feature. Returns: torch.Tensor: Random node IDs.
(x: torch.Tensor)
| 27 | |
| 28 | |
| 29 | def get_mixup_idx(x: torch.Tensor) -> torch.Tensor: |
| 30 | """ |
| 31 | Generate node IDs randomly for mixup; avoid mixup the same node. |
| 32 | |
| 33 | Args: |
| 34 | x: The latent embedding or node feature. |
| 35 | |
| 36 | Returns: |
| 37 | torch.Tensor: Random node IDs. |
| 38 | """ |
| 39 | mixup_idx = torch.randint(x.size(0) - 1, [x.size(0)]) |
| 40 | mixup_self_mask = mixup_idx - torch.arange(x.size(0)) |
| 41 | mixup_self_mask = (mixup_self_mask == 0) |
| 42 | mixup_idx += torch.ones(x.size(0), dtype=torch.int) * mixup_self_mask |
| 43 | return mixup_idx |
| 44 | |
| 45 | |
| 46 | def mixup(x: torch.Tensor, alpha: float) -> torch.Tensor: |
no outgoing calls
no test coverage detected