(g1)
| 745 | reason="python requires >= 3.7", |
| 746 | ) |
| 747 | def test_L_sym_group(g1): |
| 748 | import torch |
| 749 | |
| 750 | g1.add_hyperedges([[0, 1]], group_name="knn") |
| 751 | # all |
| 752 | H = g1.H.to_dense().cpu() |
| 753 | D_v_neg_1_2 = torch.diag(H.sum(dim=1).view(-1) ** (-0.5)) |
| 754 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 755 | W_e = g1.W_e.to_dense() |
| 756 | L_sym = ( |
| 757 | torch.eye(H.shape[0]) - D_v_neg_1_2 @ H @ W_e @ D_e_neg_1 @ H.t() @ D_v_neg_1_2 |
| 758 | ) |
| 759 | assert (L_sym == g1.L_sym.to_dense().cpu()).all() |
| 760 | # main group |
| 761 | H = g1.H_of_group("main").to_dense().cpu() |
| 762 | D_v_neg_1_2 = torch.diag(H.sum(dim=1).view(-1) ** (-0.5)) |
| 763 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 764 | W_e = g1.W_e_of_group("main").to_dense() |
| 765 | L_sym = ( |
| 766 | torch.eye(H.shape[0]) - D_v_neg_1_2 @ H @ W_e @ D_e_neg_1 @ H.t() @ D_v_neg_1_2 |
| 767 | ) |
| 768 | assert (L_sym == g1.L_sym_of_group("main").to_dense().cpu()).all() |
| 769 | # knn group |
| 770 | H = g1.H_of_group("knn").to_dense().cpu() |
| 771 | D_v_neg_1_2 = H.sum(dim=1).view(-1) ** (-0.5) |
| 772 | D_v_neg_1_2[torch.isinf(D_v_neg_1_2)] = 0 |
| 773 | D_v_neg_1_2 = torch.diag(D_v_neg_1_2) |
| 774 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 775 | W_e = g1.W_e_of_group("knn").to_dense() |
| 776 | L_sym = ( |
| 777 | torch.eye(H.shape[0]) - D_v_neg_1_2 @ H @ W_e @ D_e_neg_1 @ H.t() @ D_v_neg_1_2 |
| 778 | ) |
| 779 | assert (L_sym == g1.L_sym_of_group("knn").to_dense().cpu()).all() |
| 780 | |
| 781 | |
| 782 | @pytest.mark.skipif( |
nothing calls this directly
no test coverage detected