Generate positive and negative samples. Positive samples are generated from random walk Negative samples are generated from random sampling
(self, batch)
| 77 | self.embedding.reset_parameters() |
| 78 | |
| 79 | def sample(self, batch): |
| 80 | """ |
| 81 | Generate positive and negative samples. |
| 82 | Positive samples are generated from random walk |
| 83 | Negative samples are generated from random sampling |
| 84 | """ |
| 85 | if not isinstance(batch, torch.Tensor): |
| 86 | batch = torch.tensor(batch) |
| 87 | |
| 88 | batch = batch.repeat(self.num_walks) |
| 89 | # positive |
| 90 | pos_traces = node2vec_random_walk( |
| 91 | self.g, batch, self.p, self.q, self.walk_length, self.prob |
| 92 | ) |
| 93 | pos_traces = pos_traces.unfold(1, self.window_size, 1) # rolling window |
| 94 | pos_traces = pos_traces.contiguous().view(-1, self.window_size) |
| 95 | |
| 96 | # negative |
| 97 | neg_batch = batch.repeat(self.num_negatives) |
| 98 | neg_traces = torch.randint( |
| 99 | self.N, (neg_batch.size(0), self.walk_length) |
| 100 | ) |
| 101 | neg_traces = torch.cat([neg_batch.view(-1, 1), neg_traces], dim=-1) |
| 102 | neg_traces = neg_traces.unfold(1, self.window_size, 1) # rolling window |
| 103 | neg_traces = neg_traces.contiguous().view(-1, self.window_size) |
| 104 | |
| 105 | return pos_traces, neg_traces |
| 106 | |
| 107 | def forward(self, nodes=None): |
| 108 | """ |
nothing calls this directly
no test coverage detected