(g1)
| 681 | reason="python requires >= 3.7", |
| 682 | ) |
| 683 | def test_L_HGNN_group(g1): |
| 684 | import torch |
| 685 | |
| 686 | g1.add_hyperedges([[0, 1]], group_name="knn") |
| 687 | # all |
| 688 | H = g1.H.to_dense().cpu() |
| 689 | D_v_neg_1_2 = torch.diag(H.sum(dim=1).view(-1) ** (-0.5)) |
| 690 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 691 | W_e = g1.W_e.to_dense() |
| 692 | L_HGNN = D_v_neg_1_2 @ H @ W_e @ D_e_neg_1 @ H.t() @ D_v_neg_1_2 |
| 693 | assert (L_HGNN == g1.L_HGNN.to_dense().cpu()).all() |
| 694 | # main group |
| 695 | H = g1.H_of_group("main").to_dense().cpu() |
| 696 | D_v_neg_1_2 = torch.diag(H.sum(dim=1).view(-1) ** (-0.5)) |
| 697 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 698 | W_e = g1.W_e_of_group("main").to_dense() |
| 699 | L_HGNN = D_v_neg_1_2 @ H @ W_e @ D_e_neg_1 @ H.t() @ D_v_neg_1_2 |
| 700 | assert (L_HGNN == g1.L_HGNN_of_group("main").to_dense().cpu()).all() |
| 701 | # knn group |
| 702 | H = g1.H_of_group("knn").to_dense().cpu() |
| 703 | D_v_neg_1_2 = H.sum(dim=1).view(-1) ** (-0.5) |
| 704 | D_v_neg_1_2[torch.isinf(D_v_neg_1_2)] = 0 |
| 705 | D_v_neg_1_2 = torch.diag(D_v_neg_1_2) |
| 706 | D_e_neg_1 = torch.diag(H.sum(dim=0).view(-1) ** (-1)) |
| 707 | W_e = g1.W_e_of_group("knn").to_dense() |
| 708 | L_HGNN = D_v_neg_1_2 @ H @ W_e @ D_e_neg_1 @ H.t() @ D_v_neg_1_2 |
| 709 | assert (L_HGNN == g1.L_HGNN_of_group("knn").to_dense().cpu()).all() |
| 710 | |
| 711 | |
| 712 | @pytest.mark.skipif( |
nothing calls this directly
no test coverage detected