Team Ai
Modelpublic

kozo2/edge-ML-node2vec

sourceHugging Facecc-by-4.0updated 20d agoView on Hugging Face
0likes
node2vec_model.py116 linesDownload Raw Back to root
1"""Node2Vec model for the edge_ML_expected_ge5 graph.2 3Default run is a smoke check (builds the model, runs a few optimizer steps).4Pass --epochs N to train, which writes embeddings to --out.5"""6 7import argparse8import time9 10import torch11from torch_geometric.data import Data12from torch_geometric.data.data import DataEdgeAttr, DataTensorAttr13from torch_geometric.data.storage import BaseStorage, EdgeStorage, GlobalStorage14from torch_geometric.nn import Node2Vec15 16from paths import EMB_PATH, GRAPH_PATH17 18GRAPH = GRAPH_PATH19 20 21def load_graph(path: str = GRAPH) -> Data:22    torch.serialization.add_safe_globals(23        [Data, DataEdgeAttr, DataTensorAttr, BaseStorage, EdgeStorage, GlobalStorage]24    )25    return torch.load(path, weights_only=True)26 27 28def build_model(data: Data, args: argparse.Namespace, device: torch.device) -> Node2Vec:29    return Node2Vec(30        data.edge_index,31        embedding_dim=args.embedding_dim,32        walk_length=args.walk_length,33        context_size=args.context_size,34        walks_per_node=args.walks_per_node,35        num_negative_samples=args.num_negative_samples,36        p=args.p,37        q=args.q,38        num_nodes=data.num_nodes,39        sparse=True,  # pairs with SparseAdam; the embedding table is the only param40    ).to(device)41 42 43def main() -> None:44    ap = argparse.ArgumentParser()45    ap.add_argument("--embedding-dim", type=int, default=128)46    ap.add_argument("--walk-length", type=int, default=20)47    ap.add_argument("--context-size", type=int, default=10)48    ap.add_argument("--walks-per-node", type=int, default=10)49    ap.add_argument("--num-negative-samples", type=int, default=1)50    ap.add_argument("--p", type=float, default=1.0, help="return parameter")51    ap.add_argument("--q", type=float, default=1.0, help="in-out parameter")52    ap.add_argument("--batch-size", type=int, default=128)53    ap.add_argument("--lr", type=float, default=0.01)54    ap.add_argument("--num-workers", type=int, default=4)55    ap.add_argument("--epochs", type=int, default=0, help="0 = smoke check only")56    ap.add_argument("--steps", type=int, default=5, help="steps for the smoke check")57    ap.add_argument("--out", default=EMB_PATH)58    args = ap.parse_args()59 60    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")61    data = load_graph()62    model = build_model(data, args, device)63 64    print(f"graph      : {data.num_nodes:,} nodes, {data.edge_index.size(1) // 2:,} undirected edges")65    print(f"device     : {device}")66    print(f"model      : {model}")67    print(f"parameters : {sum(p.numel() for p in model.parameters()):,} "68          f"({data.num_nodes:,} x {args.embedding_dim})")69    print(f"walks      : length={args.walk_length} context={args.context_size} "70          f"per_node={args.walks_per_node} p={args.p} q={args.q}")71 72    loader = model.loader(batch_size=args.batch_size, shuffle=True,73                          num_workers=args.num_workers)74    optimizer = torch.optim.SparseAdam(list(model.parameters()), lr=args.lr)75    print(f"loader     : {len(loader):,} batches/epoch of {args.batch_size} seed nodes")76 77    def run_epoch(max_steps: int | None = None) -> float:78        model.train()79        total, n = 0.0, 080        for i, (pos_rw, neg_rw) in enumerate(loader):81            optimizer.zero_grad()82            loss = model.loss(pos_rw.to(device), neg_rw.to(device))83            loss.backward()84            optimizer.step()85            total, n = total + loss.item(), n + 186            if max_steps is not None and i + 1 >= max_steps:87                break88        return total / max(n, 1)89 90    if args.epochs == 0:91        t0 = time.perf_counter()92        loss = run_epoch(max_steps=args.steps)93        print(f"\nsmoke check: {args.steps} steps, mean loss {loss:.4f}, "94              f"{time.perf_counter() - t0:.1f}s")95        z = model()96        print(f"embeddings : {tuple(z.shape)} {z.dtype} on {z.device}")97        print("model built and training step verified; pass --epochs N to train")98        return99 100    for epoch in range(1, args.epochs + 1):101        t0 = time.perf_counter()102        loss = run_epoch()103        print(f"epoch {epoch:>3}/{args.epochs}  loss {loss:.4f}  "104              f"{time.perf_counter() - t0:.1f}s")105 106    model.eval()107    with torch.no_grad():108        z = model().cpu()109    torch.save({"embedding": z, "node_id": data.node_id,110                "args": vars(args)}, args.out)111    print(f"saved embeddings {tuple(z.shape)} -> {args.out}")112 113 114if __name__ == "__main__":115    main()116