(g1)
| 799 | reason="python requires >= 3.7", |
| 800 | ) |
| 801 | def test_L_rw_group(g1): |
| 802 | import torch |
| 803 | |
| 804 | g1.add_hyperedges([[0, 1]], group_name="knn") |
| 805 | # all |
| 806 | H = g1.H.to_dense().cpu() |
| 807 | D_v_neg_1 = torch.diag(H.sum(dim=1).view(-1) ** (-1)) |
| 808 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 809 | W_e = g1.W_e.to_dense() |
| 810 | L_rw = torch.eye(H.shape[0]) - D_v_neg_1 @ H @ W_e @ D_e_neg_1 @ H.t() |
| 811 | assert (L_rw == g1.L_rw.to_dense().cpu()).all() |
| 812 | # main group |
| 813 | H = g1.H_of_group("main").to_dense().cpu() |
| 814 | D_v_neg_1 = torch.diag(H.sum(dim=1).view(-1) ** (-1)) |
| 815 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 816 | W_e = g1.W_e_of_group("main").to_dense() |
| 817 | L_rw = torch.eye(H.shape[0]) - D_v_neg_1 @ H @ W_e @ D_e_neg_1 @ H.t() |
| 818 | assert (L_rw == g1.L_rw_of_group("main").to_dense().cpu()).all() |
| 819 | # knn group |
| 820 | H = g1.H_of_group("knn").to_dense().cpu() |
| 821 | D_v_neg_1 = H.sum(dim=1).view(-1) ** (-1) |
| 822 | D_v_neg_1[torch.isinf(D_v_neg_1)] = 0 |
| 823 | D_v_neg_1 = torch.diag(D_v_neg_1) |
| 824 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 825 | W_e = g1.W_e_of_group("knn").to_dense() |
| 826 | L_rw = torch.eye(H.shape[0]) - D_v_neg_1 @ H @ W_e @ D_e_neg_1 @ H.t() |
| 827 | assert (L_rw == g1.L_rw_of_group("knn").to_dense().cpu()).all() |
| 828 | |
| 829 | |
| 830 | @pytest.mark.skipif( |
nothing calls this directly
no test coverage detected