:param x: :param rate: :param noise_shape: int scalar :return:
(x, rate, noise_shape)
| 576 | |
| 577 | |
| 578 | def sparse_dropout(x, rate, noise_shape): |
| 579 | """ |
| 580 | :param x: |
| 581 | :param rate: |
| 582 | :param noise_shape: int scalar |
| 583 | :return: |
| 584 | """ |
| 585 | random_tensor = 1 - rate |
| 586 | random_tensor += torch.rand(noise_shape).to(x.device) |
| 587 | dropout_mask = torch.floor(random_tensor).byte() |
| 588 | i = x._indices() # [2, 49216] |
| 589 | v = x._values() # [49216] |
| 590 | # [2, 4926] => [49216, 2] => [remained node, 2] => [2, remained node] |
| 591 | i = i[:, dropout_mask] |
| 592 | v = v[dropout_mask] |
| 593 | out = torch.sparse.FloatTensor(i, v, x.shape).to(x.device) |
| 594 | out = out * (1./ (1-rate)) |
| 595 | return out |
| 596 | |
| 597 | |
| 598 | def dot(x, y, sparse=False): |
nothing calls this directly
no outgoing calls
no test coverage detected