viclickbait_gnn / src /build_graph.py
minhy112's picture
Upload viclickbait_gnn project
877049d verified
Raw
History Blame Contribute Delete
2.36 kB
import argparse
from pathlib import Path
import numpy as np
import torch
from sklearn.neighbors import NearestNeighbors
from utils import ensure_dir
def build_knn_graph(features: np.ndarray, k: int) -> tuple[torch.Tensor, torch.Tensor]:
neighbors = NearestNeighbors(n_neighbors=k + 1, metric="cosine")
neighbors.fit(features)
distances, indices = neighbors.kneighbors(features)
edge_map: dict[tuple[int, int], float] = {}
for node_index, (node_distances, node_neighbors) in enumerate(zip(distances, indices)):
for distance, neighbor_index in zip(node_distances[1:], node_neighbors[1:]):
weight = max(0.0, 1.0 - float(distance))
forward = (node_index, int(neighbor_index))
backward = (int(neighbor_index), node_index)
edge_map[forward] = max(edge_map.get(forward, 0.0), weight)
edge_map[backward] = max(edge_map.get(backward, 0.0), weight)
edge_items = sorted(edge_map.items())
edge_index = torch.tensor([[src, dst] for (src, dst), _ in edge_items], dtype=torch.long).t().contiguous()
edge_weight = torch.tensor([weight for _, weight in edge_items], dtype=torch.float32)
return edge_index, edge_weight
def main() -> None:
parser = argparse.ArgumentParser(description="Build a k-NN similarity graph from node features.")
parser.add_argument("--features", required=True, help="Path to node_features.npy")
parser.add_argument("--output", required=True, help="Path to graph .pt file")
parser.add_argument("--k", type=int, default=5)
parser.add_argument("--store_features", action="store_true")
args = parser.parse_args()
features = np.load(args.features)
edge_index, edge_weight = build_knn_graph(features, k=args.k)
graph = {
"edge_index": edge_index,
"edge_weight": edge_weight,
"num_nodes": int(features.shape[0]),
"feature_dim": int(features.shape[1]),
"k": args.k,
"features_path": str(Path(args.features).resolve()),
}
if args.store_features:
graph["node_features"] = torch.tensor(features, dtype=torch.float32)
output_path = Path(args.output)
ensure_dir(output_path.parent)
torch.save(graph, output_path)
print(f"Saved graph to {output_path} with {edge_index.shape[1]} directed edges")
if __name__ == "__main__":
main()