| """ |
| Tests for embed_hierarchy.py, eval_hierarchy.py, data_pbdb_taxonomy.py, |
| and synthetic_tree.py. |
| """ |
| from __future__ import annotations |
| import os |
| import sys |
|
|
| import pytest |
| import torch |
|
|
| sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) |
|
|
| from src.embed_hierarchy import HierarchyEmbedding, ranking_loss, negative_sample, train_hierarchy_embedding |
| from src.eval_hierarchy import mean_rank_and_map, compare_geometries, radius_diagnostics |
| from src.data_pbdb_taxonomy import build_edge_list, hash_edge_list, TaxonomyEdgeDataset, get_pbdb_taxonomy_dataset |
| from src.synthetic_tree import generate_synthetic_tree, get_synthetic_tree_dataset |
| from src.provenance import SchemaValidationError |
|
|
|
|
| |
| |
| |
|
|
| def test_build_edge_list_extracts_real_chain(): |
| records = [ |
| {"phylum": "Chordata", "class": "Mammalia", "order": "Carnivora", |
| "family": "Canidae", "genus": "Canis"}, |
| ] |
| edges, attrs = build_edge_list(records) |
| assert ("Mammalia", "Chordata") in edges |
| assert ("Carnivora", "Mammalia") in edges |
| assert ("Canidae", "Carnivora") in edges |
| assert ("Canis", "Canidae") in edges |
| assert attrs["Canis"]["rank"] == "genus" |
| assert attrs["Chordata"]["rank"] == "phylum" |
|
|
|
|
| def test_build_edge_list_never_fabricates_missing_link(): |
| """A record missing 'order' must NOT produce a fabricated edge |
| (family, class) skipping over the missing rank -- that would invent |
| a taxonomic claim the data never made.""" |
| records = [ |
| {"phylum": "Chordata", "class": "Mammalia", "order": None, |
| "family": "Canidae", "genus": "Canis"}, |
| {"phylum": "Chordata", "class": "Mammalia", "order": "Carnivora", |
| "family": None, "genus": None}, |
| ] |
| edges, _ = build_edge_list(records) |
| assert ("Canidae", "Mammalia") not in edges |
| assert ("Canidae", None) not in edges |
| assert ("Mammalia", "Chordata") in edges |
| assert ("Carnivora", "Mammalia") in edges |
|
|
|
|
| def test_build_edge_list_raises_on_insufficient_data(): |
| records = [{"phylum": "Chordata", "class": None, "order": None, "family": None, "genus": None}] |
| with pytest.raises(SchemaValidationError) as exc_info: |
| build_edge_list(records) |
| assert exc_info.value.outcome_code == "INSUFFICIENT_TAXONOMY_EDGES" |
|
|
|
|
| def test_build_edge_list_respects_min_rank(): |
| records = [ |
| {"phylum": "Chordata", "class": "Mammalia", "order": "Carnivora", |
| "family": "Canidae", "genus": "Canis"}, |
| {"phylum": "Chordata", "class": "Mammalia", "order": "Rodentia", |
| "family": "Muridae", "genus": "Mus"}, |
| ] |
| edges, _ = build_edge_list(records, min_rank="family") |
| assert ("Canis", "Canidae") not in edges |
| assert ("Canidae", "Carnivora") in edges |
|
|
|
|
| def test_hash_edge_list_order_independent(): |
| e1 = [("a", "b"), ("c", "d")] |
| e2 = [("c", "d"), ("a", "b")] |
| assert hash_edge_list(e1) == hash_edge_list(e2) |
|
|
|
|
| def test_hash_edge_list_differs_for_different_edges(): |
| assert hash_edge_list([("a", "b")]) != hash_edge_list([("a", "c")]) |
|
|
|
|
| |
| |
| |
|
|
| def test_taxonomy_edge_dataset_node_indices_stable(): |
| edges = [("Canis", "Canidae"), ("Canidae", "Carnivora")] |
| attrs = {"Canis": {}, "Canidae": {}, "Carnivora": {}} |
| ds = TaxonomyEdgeDataset(edges, attrs) |
| assert ds.num_nodes == 3 |
| assert len(ds) == 2 |
| item = ds[0] |
| assert isinstance(item["child_idx"], int) |
| assert isinstance(item["parent_idx"], int) |
|
|
|
|
| def test_get_pbdb_taxonomy_dataset_works_without_coordinates(tmp_path): |
| import os |
| from unittest.mock import patch, MagicMock |
|
|
| classext_only_records = [ |
| {"occurrence_no": 1001, "phylum": "Chordata", "class": "Mammalia", |
| "order": "Carnivora", "family": "Canidae", "genus": "Canis"}, |
| {"occurrence_no": 1002, "phylum": "Chordata", "class": "Mammalia", |
| "order": "Rodentia", "family": "Muridae", "genus": "Mus"}, |
| ] |
|
|
| def fake_get(url, params=None, headers=None, timeout=None): |
| resp = MagicMock() |
| resp.status_code = 200 |
| resp.json.return_value = {"records": classext_only_records} |
| return resp |
|
|
| with patch("requests.get", side_effect=fake_get): |
| dataset, meta, provenance = get_pbdb_taxonomy_dataset( |
| base_names=["TestTaxon"], |
| cache_dir=str(tmp_path / "cache"), |
| ) |
| assert provenance == "REAL_PBDB_TAXONOMY" |
| assert dataset.num_nodes > 0 |
|
|
|
|
| def test_get_pbdb_taxonomy_dataset_tolerates_genus_missing_from_first_record(tmp_path): |
| from unittest.mock import patch, MagicMock |
|
|
| records_genus_sparse = [ |
| {"occurrence_no": 1, "phylum": "Chordata", "class": "Mammalia", |
| "order": "Carnivora", "family": "Canidae"}, |
| {"occurrence_no": 2, "phylum": "Chordata", "class": "Mammalia", |
| "order": "Carnivora", "family": "Canidae", "genus": "Canis"}, |
| ] |
|
|
| def fake_get(url, params=None, headers=None, timeout=None): |
| resp = MagicMock() |
| resp.status_code = 200 |
| resp.json.return_value = {"records": records_genus_sparse} |
| return resp |
|
|
| with patch("requests.get", side_effect=fake_get): |
| dataset, meta, provenance = get_pbdb_taxonomy_dataset( |
| base_names=["TestTaxon"], |
| cache_dir=str(tmp_path / "cache2"), |
| ) |
| assert provenance == "REAL_PBDB_TAXONOMY" |
| assert dataset.num_nodes > 0 |
|
|
|
|
| |
| |
| |
|
|
| def test_synthetic_tree_is_connected_and_acyclic(): |
| edges, attrs = generate_synthetic_tree(n_nodes=50, seed=1) |
| assert len(edges) == 49 |
| assert len(attrs) == 50 |
| |
| for child, parent in edges: |
| assert attrs[child]["depth"] == attrs[parent]["depth"] + 1 |
|
|
|
|
| def test_synthetic_tree_deterministic_given_seed(): |
| e1, _ = generate_synthetic_tree(n_nodes=30, seed=42) |
| e2, _ = generate_synthetic_tree(n_nodes=30, seed=42) |
| assert e1 == e2 |
|
|
|
|
| def test_balanced_tree_shape_and_depth(): |
| from src.synthetic_tree import generate_balanced_tree |
| edges, attrs = generate_balanced_tree(branching_factor=3, depth=4) |
| |
| assert len(attrs) == 1 + 3 + 9 + 27 + 81 |
| assert len(edges) == len(attrs) - 1 |
| assert max(a["depth"] for a in attrs.values()) == 4 |
| for child, parent in edges: |
| assert attrs[child]["depth"] == attrs[parent]["depth"] + 1 |
|
|
|
|
| def test_synthetic_tree_dataset_provenance(): |
| ds, meta, provenance = get_synthetic_tree_dataset(n_nodes=20, seed=0) |
| assert provenance == "SYNTHETIC_TREE" |
| assert ds.num_nodes == 20 |
|
|
|
|
| |
| |
| |
|
|
| def test_hierarchy_embedding_rejects_invalid_geometry(): |
| with pytest.raises(ValueError): |
| HierarchyEmbedding(num_nodes=5, dim=4, geometry="spherical") |
|
|
|
|
| def test_hierarchy_embedding_points_shape_both_geometries(): |
| for geom in ("poincare", "euclidean"): |
| model = HierarchyEmbedding(num_nodes=10, dim=4, geometry=geom) |
| idx = torch.tensor([0, 1, 2]) |
| pts = model.points(idx) |
| assert pts.shape == (3, 4) |
|
|
|
|
| def test_hierarchy_embedding_poincare_points_stay_in_ball(): |
| model = HierarchyEmbedding(num_nodes=20, dim=4, geometry="poincare", init_scale=1e-3) |
| idx = torch.arange(20) |
| pts = model.points(idx) |
| assert (pts.norm(dim=-1) < 1.0).all() |
|
|
|
|
| def test_hierarchy_embedding_distance_symmetric_both_geometries(): |
| for geom in ("poincare", "euclidean"): |
| model = HierarchyEmbedding(num_nodes=10, dim=4, geometry=geom) |
| a = model.points(torch.tensor([0, 1])) |
| b = model.points(torch.tensor([2, 3])) |
| d_ab = model.distance(a, b) |
| d_ba = model.distance(b, a) |
| assert torch.allclose(d_ab, d_ba, atol=1e-5) |
|
|
|
|
| def test_hierarchy_embedding_optimizer_actually_updates_params(): |
| for geom in ("poincare", "euclidean"): |
| model = HierarchyEmbedding(num_nodes=10, dim=4, geometry=geom) |
| before = model.emb.detach().clone() |
| opt = model.make_optimizer(lr=0.1) |
| loss = model.points(torch.arange(10)).pow(2).sum() |
| opt.zero_grad() |
| loss.backward() |
| opt.step() |
| after = model.emb.detach().clone() |
| assert not torch.allclose(before, after), f"{geom} params did not move" |
|
|
|
|
| |
| |
| |
|
|
| def test_negative_sample_never_equals_true_parent(): |
| child_idx = torch.tensor([0, 1, 2, 3, 4]) |
| true_parent = torch.tensor([5, 5, 5, 5, 5]) |
| neg = negative_sample(child_idx, true_parent, num_nodes=6, k=20) |
| assert not (neg == true_parent.unsqueeze(1)).any() |
|
|
|
|
| def test_negative_sample_never_equals_child_itself(): |
| child_idx = torch.tensor([0, 1, 2, 3, 4]) |
| true_parent = torch.tensor([5, 5, 5, 5, 5]) |
| neg = negative_sample(child_idx, true_parent, num_nodes=6, k=20) |
| assert not (neg == child_idx.unsqueeze(1)).any() |
|
|
|
|
| |
| |
| |
|
|
| def test_ranking_loss_is_zero_when_parent_much_closer_than_negative(): |
| model = HierarchyEmbedding(num_nodes=3, dim=2, geometry="euclidean") |
| with torch.no_grad(): |
| model.emb[0] = torch.tensor([0.0, 0.0]) |
| model.emb[1] = torch.tensor([0.01, 0.0]) |
| model.emb[2] = torch.tensor([10.0, 10.0]) |
| loss = ranking_loss( |
| model, torch.tensor([0]), torch.tensor([1]), torch.tensor([[2]]), margin=1.0, |
| ) |
| assert loss.item() == pytest.approx(0.0, abs=1e-4) |
|
|
|
|
| def test_ranking_loss_is_positive_when_negative_closer_than_parent(): |
| model = HierarchyEmbedding(num_nodes=3, dim=2, geometry="euclidean") |
| with torch.no_grad(): |
| model.emb[0] = torch.tensor([0.0, 0.0]) |
| model.emb[1] = torch.tensor([10.0, 10.0]) |
| model.emb[2] = torch.tensor([0.01, 0.0]) |
| loss = ranking_loss( |
| model, torch.tensor([0]), torch.tensor([1]), torch.tensor([[2]]), margin=1.0, |
| ) |
| assert loss.item() > 0.5 |
|
|
|
|
| |
| |
| |
|
|
| def test_mean_rank_and_map_perfect_recovery_gives_rank_one(): |
| model = HierarchyEmbedding(num_nodes=4, dim=2, geometry="euclidean") |
| with torch.no_grad(): |
| model.emb[0] = torch.tensor([0.0, 0.0]) |
| model.emb[1] = torch.tensor([0.1, 0.0]) |
| model.emb[2] = torch.tensor([5.0, 0.0]) |
| model.emb[3] = torch.tensor([9.0, 0.0]) |
| metrics = mean_rank_and_map(model, edges=[(0, 1)]) |
| assert metrics["mean_rank"] == pytest.approx(1.0) |
| assert metrics["mrr"] == pytest.approx(1.0) |
| assert metrics["hits@1"] == pytest.approx(1.0) |
|
|
|
|
| def test_mean_rank_and_map_worst_case_gives_high_rank(): |
| model = HierarchyEmbedding(num_nodes=4, dim=2, geometry="euclidean") |
| with torch.no_grad(): |
| model.emb[0] = torch.tensor([0.0, 0.0]) |
| model.emb[1] = torch.tensor([9.0, 0.0]) |
| model.emb[2] = torch.tensor([0.1, 0.0]) |
| model.emb[3] = torch.tensor([0.2, 0.0]) |
| metrics = mean_rank_and_map(model, edges=[(0, 1)]) |
| assert metrics["mean_rank"] == pytest.approx(3.0) |
| assert metrics["hits@1"] == 0.0 |
|
|
|
|
| def test_compare_geometries_returns_deltas(): |
| p_model = HierarchyEmbedding(num_nodes=5, dim=2, geometry="poincare") |
| e_model = HierarchyEmbedding(num_nodes=5, dim=2, geometry="euclidean") |
| result = compare_geometries(p_model, e_model, test_edges=[(0, 1), (1, 2)]) |
| assert "delta_mrr" in result |
| assert "poincare" in result and "euclidean" in result |
|
|
|
|
| def test_radius_diagnostics_shape(): |
| model = HierarchyEmbedding(num_nodes=5, dim=3, geometry="poincare") |
| depths = {i: i for i in range(5)} |
| result = radius_diagnostics(model, depths) |
| assert "radius_depth_correlation" in result |
| assert -1.0 <= result["radius_depth_correlation"] <= 1.0 + 1e-6 |
|
|
|
|
| |
| |
| |
|
|
| def test_train_hierarchy_embedding_loss_decreases_both_geometries(): |
| ds, meta, provenance = get_synthetic_tree_dataset(n_nodes=40, seed=7) |
| for geom in ("poincare", "euclidean"): |
| model, metrics = train_hierarchy_embedding( |
| ds, geometry=geom, dim=4, epochs=15, batch_size=16, |
| lr=0.05, neg_samples=5, seed=0, |
| ) |
| history = metrics["loss_history"] |
| assert history[-1] < history[0], f"{geom}: loss did not decrease ({history[0]} -> {history[-1]})" |
|
|
|
|
| def test_softmax_ranking_loss_penalizes_correctly(): |
| from src.embed_hierarchy import softmax_ranking_loss |
| model = HierarchyEmbedding(num_nodes=3, dim=2, geometry="euclidean") |
| with torch.no_grad(): |
| model.emb[0] = torch.tensor([0.0, 0.0]) |
| model.emb[1] = torch.tensor([0.01, 0.0]) |
| model.emb[2] = torch.tensor([10.0, 10.0]) |
| loss_good = softmax_ranking_loss( |
| model, torch.tensor([0]), torch.tensor([1]), torch.tensor([[2]]), |
| ) |
| with torch.no_grad(): |
| model.emb[1] = torch.tensor([10.0, 10.0]) |
| model.emb[2] = torch.tensor([0.01, 0.0]) |
| loss_bad = softmax_ranking_loss( |
| model, torch.tensor([0]), torch.tensor([1]), torch.tensor([[2]]), |
| ) |
| assert loss_good.item() < loss_bad.item() |
| assert loss_good.item() < 0.01 |
|
|
|
|
| def test_train_hierarchy_embedding_softmax_loss_type_runs_and_decreases(): |
| ds, meta, provenance = get_synthetic_tree_dataset(n_nodes=40, seed=7) |
| model, metrics = train_hierarchy_embedding( |
| ds, geometry="poincare", dim=4, epochs=15, batch_size=16, |
| lr=0.05, neg_samples=5, seed=0, loss_type="softmax", |
| ) |
| history = metrics["loss_history"] |
| assert history[-1] < history[0] |
|
|
|
|
| def test_train_hierarchy_embedding_rejects_invalid_loss_type(): |
| ds, meta, provenance = get_synthetic_tree_dataset(n_nodes=20, seed=1) |
| with pytest.raises(ValueError): |
| train_hierarchy_embedding(ds, dim=2, epochs=1, loss_type="not_a_real_loss") |
|
|
|
|
| def test_train_hierarchy_embedding_burn_in_restores_lr_after_burn_in(): |
| ds, meta, provenance = get_synthetic_tree_dataset(n_nodes=20, seed=1) |
| model, metrics = train_hierarchy_embedding( |
| ds, geometry="euclidean", dim=4, epochs=10, batch_size=8, |
| lr=0.05, neg_samples=5, seed=0, burn_in_epochs=3, burn_in_lr_mult=0.1, |
| ) |
| assert all(torch.isfinite(torch.tensor(x)) for x in metrics["loss_history"]) |
| assert len(metrics["loss_history"]) == 10 |
|
|
|
|
| def test_learnable_curvature_does_not_diverge_on_the_scenario_that_broke_it(): |
| ds, meta, provenance = get_synthetic_tree_dataset( |
| n_nodes=200, seed=0, tree_type="balanced", branching_factor=3, depth=5, |
| ) |
| c_min, c_max = 0.1, 3.0 |
| model, metrics = train_hierarchy_embedding( |
| ds, geometry="poincare", dim=2, epochs=60, batch_size=64, |
| lr=0.02, neg_samples=10, seed=0, loss_type="softmax", |
| burn_in_epochs=10, burn_in_lr_mult=0.1, learnable_c=True, |
| curvature_lr_mult=0.1, c_min=c_min, c_max=c_max, |
| ) |
|
|
| assert all(torch.isfinite(torch.tensor(x)) for x in metrics["loss_history"]), \ |
| "loss went non-finite at some epoch" |
|
|
| for epoch_i, c_val in enumerate(metrics["c_history"]): |
| assert c_min - 1e-6 <= c_val <= c_max + 1e-6, \ |
| f"epoch {epoch_i}: curvature {c_val} escaped [{c_min}, {c_max}]" |
|
|
| for epoch_i, (c_val, max_norm) in enumerate(zip(metrics["c_history"], metrics["max_norm_history"])): |
| current_radius = 0.99 / (c_val ** 0.5) |
| assert max_norm < current_radius, \ |
| (f"epoch {epoch_i}: max point norm {max_norm:.4f} outside the CURRENT " |
| f"ball radius {current_radius:.4f} for c={c_val:.4f} -- points are " |
| f"invalid under the curvature they're actually being evaluated at") |
|
|
|
|
| def test_clip_to_ball_keeps_points_strictly_inside_current_curvature_radius(): |
| model = HierarchyEmbedding(num_nodes=10, dim=4, geometry="poincare", init_scale=1e-3) |
| with torch.no_grad(): |
| model.emb.data.mul_(10000.0) |
| model.clip_to_ball(margin=0.95) |
| norms = model.emb.detach().norm(dim=-1) |
| current_radius = 1.0 / model.manifold.c.clamp_min(1e-8).sqrt() |
| assert (norms <= current_radius + 1e-6).all() |
|
|
|
|
| def test_clip_to_ball_is_noop_for_euclidean(): |
| model = HierarchyEmbedding(num_nodes=10, dim=4, geometry="euclidean") |
| before = model.emb.detach().clone() |
| model.clip_to_ball() |
| after = model.emb.detach().clone() |
| assert torch.allclose(before, after) |
|
|
|
|
| def test_clamp_curvature_actually_constrains_the_leaf_not_a_derived_view(): |
| model = HierarchyEmbedding(num_nodes=5, dim=2, geometry="poincare", learnable_c=True, c=1.0) |
| with torch.no_grad(): |
| model.manifold.isp_c.fill_(100.0) |
| model.clamp_curvature(c_min=0.1, c_max=3.0) |
| assert model.manifold.c.item() <= 3.0 + 1e-4 |
|
|
| with torch.no_grad(): |
| model.manifold.isp_c.fill_(-100.0) |
| model.clamp_curvature(c_min=0.1, c_max=3.0) |
| assert model.manifold.c.item() >= 0.1 - 1e-4 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(pytest.main([__file__, "-v"])) |
|
|