File size: 3,059 Bytes
ae73c7f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
"""
Synthetic random tree, for validating the embedding/training/evaluation
pipeline WITHOUT needing live PBDB network access (blocked in this
development sandbox). Provenance is explicitly "SYNTHETIC_TREE" — never
to be confused with real PBDB taxonomy data, and never silently
substituted for it (same discipline as get_synthetic_dataset() in
data_real.py: explicit opt-in only, clearly labeled).

Construction: a random recursive tree — node i (for i >= 1) attaches to
a uniformly random existing node in {0, ..., i-1}. This guarantees a
single connected tree with no cycles, node 0 as root, and is a standard,
well-understood random tree model (not an arbitrary ad hoc generator).
"""
from __future__ import annotations
from typing import Dict, List, Tuple

import torch

from .data_pbdb_taxonomy import TaxonomyEdgeDataset


def generate_synthetic_tree(
    n_nodes: int = 300,
    seed: int = 0,
) -> Tuple[List[Tuple[str, str]], Dict[str, Dict[str, int]]]:
    gen = torch.Generator().manual_seed(seed)
    edges: List[Tuple[str, str]] = []
    node_attrs: Dict[str, Dict[str, int]] = {"n0": {"depth": 0}}
    depth = [0] * n_nodes
    for i in range(1, n_nodes):
        parent = int(torch.randint(0, i, (1,), generator=gen).item())
        edges.append((f"n{i}", f"n{parent}"))
        depth[i] = depth[parent] + 1
        node_attrs[f"n{i}"] = {"depth": depth[i]}
    return edges, node_attrs


def generate_balanced_tree(
    branching_factor: int = 3,
    depth: int = 6,
) -> Tuple[List[Tuple[str, str]], Dict[str, Dict[str, int]]]:

    edges: List[Tuple[str, str]] = []
    node_attrs: Dict[str, Dict[str, int]] = {"n0": {"depth": 0}}
    frontier = ["n0"]
    counter = 1
    for d in range(1, depth + 1):
        next_frontier = []
        for parent in frontier:
            for _ in range(branching_factor):
                name = f"n{counter}"
                counter += 1
                edges.append((name, parent))
                node_attrs[name] = {"depth": d}
                next_frontier.append(name)
        frontier = next_frontier
    return edges, node_attrs


def get_synthetic_tree_dataset(n_nodes: int = 300, seed: int = 0, tree_type: str = "random_recursive",
                                 branching_factor: int = 3, depth: int = 6):
    if tree_type == "balanced":
        edges, node_attrs = generate_balanced_tree(branching_factor=branching_factor, depth=depth)
    else:
        edges, node_attrs = generate_synthetic_tree(n_nodes=n_nodes, seed=seed)
    dataset = TaxonomyEdgeDataset(edges, node_attrs)
    node_depths_by_idx = {
        dataset.node_to_idx[name]: attrs["depth"] for name, attrs in node_attrs.items()
    }
    meta = {
        "num_nodes": dataset.num_nodes,
        "num_edges": len(edges),
        "seed": seed,
        "tree_type": tree_type,
        "node_depths_by_idx": node_depths_by_idx,
    }
    print(f"[data] SYNTHETIC_TREE ({tree_type}): {dataset.num_nodes} nodes, {len(edges)} edges "
          f"(seed={seed}) — NOT real PBDB data")
    return dataset, meta, "SYNTHETIC_TREE"