diff --git a/fedgraph/data_process.py b/fedgraph/data_process.py index 803995d..b1a9d1c 100644 --- a/fedgraph/data_process.py +++ b/fedgraph/data_process.py @@ -101,6 +101,7 @@ def data_loader_NC(args: attridict) -> tuple: ( communicate_node_global_indexes, in_com_train_node_local_indexes, + in_com_val_node_local_indexes, in_com_test_node_local_indexes, global_edge_indexes_clients, ) = get_in_comm_indexes( @@ -110,17 +111,20 @@ def data_loader_NC(args: attridict) -> tuple: args.num_hops, idx_train, idx_test, + idx_val=idx_val, ) return ( edge_index, features, labels, idx_train, + idx_val, idx_test, class_num, split_node_indexes, communicate_node_global_indexes, in_com_train_node_local_indexes, + in_com_val_node_local_indexes, in_com_test_node_local_indexes, global_edge_indexes_clients, ) diff --git a/fedgraph/federated_methods.py b/fedgraph/federated_methods.py index 64360b5..fb7bdef 100644 --- a/fedgraph/federated_methods.py +++ b/fedgraph/federated_methods.py @@ -132,6 +132,78 @@ def _validate_nc_num_hops(args: Any) -> None: ) +def _unpack_nc_data(data: tuple) -> dict: + """Return named NC data fields, supporting legacy tuples without val data.""" + if len(data) == 13: + ( + edge_index, + features, + labels, + idx_train, + idx_val, + idx_test, + class_num, + split_node_indexes, + communicate_node_global_indexes, + in_com_train_node_local_indexes, + in_com_val_node_local_indexes, + in_com_test_node_local_indexes, + global_edge_indexes_clients, + ) = data + elif len(data) == 11: + ( + edge_index, + features, + labels, + idx_train, + idx_test, + class_num, + split_node_indexes, + communicate_node_global_indexes, + in_com_train_node_local_indexes, + in_com_test_node_local_indexes, + global_edge_indexes_clients, + ) = data + idx_val = torch.empty(0, dtype=idx_train.dtype) + if isinstance(in_com_train_node_local_indexes, dict): + in_com_val_node_local_indexes = { + key: torch.empty(0, dtype=idx_train.dtype) + for key in in_com_train_node_local_indexes + } + else: + in_com_val_node_local_indexes = [ + torch.empty(0, dtype=idx_train.dtype) + for _ in in_com_train_node_local_indexes + ] + else: + raise ValueError( + "Unexpected NC data tuple format; expected 11 legacy fields or " + "13 fields with validation data." + ) + + return { + "edge_index": edge_index, + "features": features, + "labels": labels, + "idx_train": idx_train, + "idx_val": idx_val, + "idx_test": idx_test, + "class_num": class_num, + "split_node_indexes": split_node_indexes, + "communicate_node_global_indexes": communicate_node_global_indexes, + "in_com_train_node_local_indexes": in_com_train_node_local_indexes, + "in_com_val_node_local_indexes": in_com_val_node_local_indexes, + "in_com_test_node_local_indexes": in_com_test_node_local_indexes, + "global_edge_indexes_clients": global_edge_indexes_clients, + } + + +def _weighted_nc_metric(results: np.ndarray, weights: list, metric_index: int) -> float: + if not weights or sum(weights) == 0: + return 0.0 + return float(np.average([row[metric_index] for row in results], weights=weights)) + + def run_fedgraph(args: attridict) -> None: """ Run the training process for the specified task. @@ -302,19 +374,20 @@ def run_NC(args: attridict, data: Any = None) -> None: print("Changing method to FedAvg") args.method = "FedAvg" if not args.use_huggingface: - ( - edge_index, - features, - labels, - idx_train, - idx_test, - class_num, - split_node_indexes, - communicate_node_global_indexes, - in_com_train_node_local_indexes, - in_com_test_node_local_indexes, - global_edge_indexes_clients, - ) = data + nc_data = _unpack_nc_data(data) + edge_index = nc_data["edge_index"] + features = nc_data["features"] + labels = nc_data["labels"] + idx_train = nc_data["idx_train"] + idx_val = nc_data["idx_val"] + idx_test = nc_data["idx_test"] + class_num = nc_data["class_num"] + split_node_indexes = nc_data["split_node_indexes"] + communicate_node_global_indexes = nc_data["communicate_node_global_indexes"] + in_com_train_node_local_indexes = nc_data["in_com_train_node_local_indexes"] + in_com_val_node_local_indexes = nc_data["in_com_val_node_local_indexes"] + in_com_test_node_local_indexes = nc_data["in_com_test_node_local_indexes"] + global_edge_indexes_clients = nc_data["global_edge_indexes_clients"] if args.saveto_huggingface: save_all_trainers_data( split_node_indexes=split_node_indexes, @@ -323,6 +396,7 @@ def run_NC(args: attridict, data: Any = None) -> None: labels=labels, features=features, in_com_train_node_local_indexes=in_com_train_node_local_indexes, + in_com_val_node_local_indexes=in_com_val_node_local_indexes, in_com_test_node_local_indexes=in_com_test_node_local_indexes, n_trainer=args.n_trainer, class_num=class_num, @@ -424,11 +498,15 @@ def get_memory_usage(self): train_labels=labels[communicate_node_global_indexes[i]][ in_com_train_node_local_indexes[i] ], + val_labels=labels[communicate_node_global_indexes[i]][ + in_com_val_node_local_indexes[i] + ], test_labels=labels[communicate_node_global_indexes[i]][ in_com_test_node_local_indexes[i] ], features=features[split_node_indexes[i]], idx_train=in_com_train_node_local_indexes[i], + idx_val=in_com_val_node_local_indexes[i], idx_test=in_com_test_node_local_indexes[i], global_node_num=len(features), class_num=class_num, @@ -456,6 +534,9 @@ def get_memory_usage(self): train_data_weights = [ info["len_in_com_train_node_local_indexes"] for info in trainer_information ] + val_data_weights = [ + info["len_in_com_val_node_local_indexes"] for info in trainer_information + ] test_data_weights = [ info["len_in_com_test_node_local_indexes"] for info in trainer_information ] @@ -937,15 +1018,16 @@ def get_memory_usage(self): round_comm_time = comm_end - comm_start total_communication_time += round_comm_time - # Testing phase (not counted in training or communication time) - results = [trainer.local_test.remote() for trainer in server.trainers] + # Validation phase (not counted in pure training or communication time). + # Test metrics are intentionally reserved for final evaluation. + results = [trainer.local_val.remote() for trainer in server.trainers] results = np.array([ray.get(result) for result in results]) - average_test_accuracy = np.average( - [row[1] for row in results], weights=test_data_weights, axis=0 - ) - global_acc_list.append(average_test_accuracy) + average_val_loss = _weighted_nc_metric(results, val_data_weights, 0) + average_val_accuracy = _weighted_nc_metric(results, val_data_weights, 1) + global_acc_list.append(average_val_accuracy) - print(f"Round {i+1}: Global Test Accuracy = {average_test_accuracy:.4f}") + print(f"Round {i+1}: Global Val Loss = {average_val_loss:.4f}") + print(f"Round {i+1}: Global Val Accuracy = {average_val_accuracy:.4f}") print( f"Round {i+1}: Training Time = {round_training_time:.2f}s, Communication Time = {round_comm_time:.2f}s" ) @@ -1218,19 +1300,20 @@ def run_NC_dp(args: attridict, data: Any = None) -> None: args.method = "FedAvg" if not args.use_huggingface: - ( - edge_index, - features, - labels, - idx_train, - idx_test, - class_num, - split_node_indexes, - communicate_node_global_indexes, - in_com_train_node_local_indexes, - in_com_test_node_local_indexes, - global_edge_indexes_clients, - ) = data + nc_data = _unpack_nc_data(data) + edge_index = nc_data["edge_index"] + features = nc_data["features"] + labels = nc_data["labels"] + idx_train = nc_data["idx_train"] + idx_val = nc_data["idx_val"] + idx_test = nc_data["idx_test"] + class_num = nc_data["class_num"] + split_node_indexes = nc_data["split_node_indexes"] + communicate_node_global_indexes = nc_data["communicate_node_global_indexes"] + in_com_train_node_local_indexes = nc_data["in_com_train_node_local_indexes"] + in_com_val_node_local_indexes = nc_data["in_com_val_node_local_indexes"] + in_com_test_node_local_indexes = nc_data["in_com_test_node_local_indexes"] + global_edge_indexes_clients = nc_data["global_edge_indexes_clients"] if args.dataset in ["simulate", "cora", "citeseer", "pubmed", "reddit"]: args_hidden = 16 @@ -1279,11 +1362,15 @@ def __init__(self, *args: Any, **kwds: Any): train_labels=labels[communicate_node_global_indexes[i]][ in_com_train_node_local_indexes[i] ], + val_labels=labels[communicate_node_global_indexes[i]][ + in_com_val_node_local_indexes[i] + ], test_labels=labels[communicate_node_global_indexes[i]][ in_com_test_node_local_indexes[i] ], features=features[split_node_indexes[i]], idx_train=in_com_train_node_local_indexes[i], + idx_val=in_com_val_node_local_indexes[i], idx_test=in_com_test_node_local_indexes[i], global_node_num=len(features), class_num=class_num, @@ -1310,6 +1397,9 @@ def __init__(self, *args: Any, **kwds: Any): train_data_weights = [ info["len_in_com_train_node_local_indexes"] for info in trainer_information ] + val_data_weights = [ + info["len_in_com_val_node_local_indexes"] for info in trainer_information + ] test_data_weights = [ info["len_in_com_test_node_local_indexes"] for info in trainer_information ] @@ -1392,14 +1482,14 @@ def __init__(self, *args: Any, **kwds: Any): for i in range(args.global_rounds): server.train(i) - results = [trainer.local_test.remote() for trainer in server.trainers] + results = [trainer.local_val.remote() for trainer in server.trainers] results = np.array([ray.get(result) for result in results]) - average_test_accuracy = np.average( - [row[1] for row in results], weights=test_data_weights, axis=0 - ) - global_acc_list.append(average_test_accuracy) + average_val_loss = _weighted_nc_metric(results, val_data_weights, 0) + average_val_accuracy = _weighted_nc_metric(results, val_data_weights, 1) + global_acc_list.append(average_val_accuracy) - print(f"Round {i+1}: Global Test Accuracy = {average_test_accuracy:.4f}") + print(f"Round {i+1}: Global Val Loss = {average_val_loss:.4f}") + print(f"Round {i+1}: Global Val Accuracy = {average_val_accuracy:.4f}") model_size_mb = server.get_model_size() / (1024 * 1024) monitor.add_train_comm_cost( @@ -1465,19 +1555,20 @@ def run_NC_lowrank(args: attridict, data: Any = None) -> None: args.method = "FedAvg" if not args.use_huggingface: - ( - edge_index, - features, - labels, - idx_train, - idx_test, - class_num, - split_node_indexes, - communicate_node_global_indexes, - in_com_train_node_local_indexes, - in_com_test_node_local_indexes, - global_edge_indexes_clients, - ) = data + nc_data = _unpack_nc_data(data) + edge_index = nc_data["edge_index"] + features = nc_data["features"] + labels = nc_data["labels"] + idx_train = nc_data["idx_train"] + idx_val = nc_data["idx_val"] + idx_test = nc_data["idx_test"] + class_num = nc_data["class_num"] + split_node_indexes = nc_data["split_node_indexes"] + communicate_node_global_indexes = nc_data["communicate_node_global_indexes"] + in_com_train_node_local_indexes = nc_data["in_com_train_node_local_indexes"] + in_com_val_node_local_indexes = nc_data["in_com_val_node_local_indexes"] + in_com_test_node_local_indexes = nc_data["in_com_test_node_local_indexes"] + global_edge_indexes_clients = nc_data["global_edge_indexes_clients"] if args.saveto_huggingface: save_all_trainers_data( @@ -1487,6 +1578,7 @@ def run_NC_lowrank(args: attridict, data: Any = None) -> None: labels=labels, features=features, in_com_train_node_local_indexes=in_com_train_node_local_indexes, + in_com_val_node_local_indexes=in_com_val_node_local_indexes, in_com_test_node_local_indexes=in_com_test_node_local_indexes, n_trainer=args.n_trainer, class_num=class_num, @@ -1541,11 +1633,15 @@ def __init__(self, *args: Any, **kwds: Any): train_labels=labels[communicate_node_global_indexes[i]][ in_com_train_node_local_indexes[i] ], + val_labels=labels[communicate_node_global_indexes[i]][ + in_com_val_node_local_indexes[i] + ], test_labels=labels[communicate_node_global_indexes[i]][ in_com_test_node_local_indexes[i] ], features=features[split_node_indexes[i]], idx_train=in_com_train_node_local_indexes[i], + idx_val=in_com_val_node_local_indexes[i], idx_test=in_com_test_node_local_indexes[i], global_node_num=len(features), class_num=class_num, @@ -1572,6 +1668,9 @@ def __init__(self, *args: Any, **kwds: Any): train_data_weights = [ info["len_in_com_train_node_local_indexes"] for info in trainer_information ] + val_data_weights = [ + info["len_in_com_val_node_local_indexes"] for info in trainer_information + ] test_data_weights = [ info["len_in_com_test_node_local_indexes"] for info in trainer_information ] @@ -1602,15 +1701,15 @@ def __init__(self, *args: Any, **kwds: Any): for i in range(args.global_rounds): server.train(i) - # Evaluation - results = [trainer.local_test.remote() for trainer in server.trainers] + # Validation evaluation. Test metrics are reserved for final reporting. + results = [trainer.local_val.remote() for trainer in server.trainers] results = np.array([ray.get(result) for result in results]) - average_test_accuracy = np.average( - [row[1] for row in results], weights=test_data_weights, axis=0 - ) - global_acc_list.append(average_test_accuracy) + average_val_loss = _weighted_nc_metric(results, val_data_weights, 0) + average_val_accuracy = _weighted_nc_metric(results, val_data_weights, 1) + global_acc_list.append(average_val_accuracy) - print(f"Round {i+1}: Global Test Accuracy = {average_test_accuracy:.4f}") + print(f"Round {i+1}: Global Val Loss = {average_val_loss:.4f}") + print(f"Round {i+1}: Global Val Accuracy = {average_val_accuracy:.4f}") # Communication cost tracking (enhanced with compression-aware sizing) model_size_mb = server.get_model_size() / (1024 * 1024) diff --git a/fedgraph/trainer_class.py b/fedgraph/trainer_class.py index 9ff9564..dcd7ff1 100644 --- a/fedgraph/trainer_class.py +++ b/fedgraph/trainer_class.py @@ -76,10 +76,19 @@ def download_and_load_tensor(file_name, optional=False): ) global_edge_index_client = download_and_load_tensor("adj.pt") train_labels = download_and_load_tensor("train_labels.pt") + val_labels = download_and_load_tensor("val_labels.pt", optional=True) test_labels = download_and_load_tensor("test_labels.pt") features = download_and_load_tensor("features.pt") in_com_train_node_local_indexes = download_and_load_tensor("idx_train.pt") + in_com_val_node_local_indexes = download_and_load_tensor( + "idx_val.pt", optional=True + ) in_com_test_node_local_indexes = download_and_load_tensor("idx_test.pt") + if val_labels is None or in_com_val_node_local_indexes is None: + val_labels = torch.empty(0, dtype=train_labels.dtype) + in_com_val_node_local_indexes = torch.empty( + 0, dtype=in_com_train_node_local_indexes.dtype + ) global_node_num = download_and_load_tensor("global_node_num.pt", optional=True) class_num = download_and_load_tensor("class_num.pt", optional=True) if global_node_num is None or class_num is None: @@ -93,9 +102,11 @@ def download_and_load_tensor(file_name, optional=False): communicate_node_global_index, global_edge_index_client, train_labels, + val_labels, test_labels, features, in_com_train_node_local_indexes, + in_com_val_node_local_indexes, in_com_test_node_local_indexes, global_node_num, class_num, @@ -160,9 +171,11 @@ def __init__( communicate_node_index: torch.Tensor = None, adj: torch.Tensor = None, train_labels: torch.Tensor = None, + val_labels: torch.Tensor = None, test_labels: torch.Tensor = None, features: torch.Tensor = None, idx_train: torch.Tensor = None, + idx_val: torch.Tensor = None, idx_test: torch.Tensor = None, global_node_num: Optional[int] = None, class_num: Optional[int] = None, @@ -194,13 +207,19 @@ def __init__( communicate_node_index, adj, train_labels, + val_labels, test_labels, features, idx_train, + idx_val, idx_test, global_node_num, class_num, ) = load_trainer_data_from_hugging_face(rank, args) + if val_labels is None: + val_labels = torch.empty(0, dtype=train_labels.dtype) + if idx_val is None: + idx_val = torch.empty(0, dtype=idx_train.dtype) self.rank = rank # rank = trainer ID self.device = device @@ -212,15 +231,19 @@ def __init__( self.test_losses: list = [] self.test_accs: list = [] + self.val_losses: list = [] + self.val_accs: list = [] self.local_node_index = local_node_index.to(device) self.communicate_node_index = communicate_node_index.to(device) self.adj = adj.to(device) self.train_labels = train_labels.to(device) + self.val_labels = val_labels.to(device) self.test_labels = test_labels.to(device) self.features = features.to(device) self.idx_train = idx_train.to(device) + self.idx_val = idx_val.to(device) self.idx_test = idx_test.to(device) self.local_step = args.local_step @@ -246,7 +269,7 @@ def __init__( def get_info(self): label_nums = [ int(labels.max().item()) + 1 - for labels in (self.train_labels, self.test_labels) + for labels in (self.train_labels, self.val_labels, self.test_labels) if labels.numel() > 0 ] return { @@ -256,6 +279,7 @@ def get_info(self): "class_num": self.class_num, "feature_shape": self.features.shape[1], "len_in_com_train_node_local_indexes": len(self.idx_train), + "len_in_com_val_node_local_indexes": len(self.idx_val), "len_in_com_test_node_local_indexes": len(self.idx_test), "communicate_node_global_index": self.communicate_node_index, } @@ -829,7 +853,8 @@ def train(self, current_global_round: int) -> None: self.feature_aggregation = self.feature_aggregation.to(self.device) data = None - if hasattr(self.args, "batch_size") and self.args.batch_size > 0: + use_mini_batch = hasattr(self.args, "batch_size") and self.args.batch_size > 0 + if use_mini_batch: # batch preparation train_mask = torch.zeros( self.feature_aggregation.size(0), dtype=torch.bool @@ -852,40 +877,43 @@ def train(self, current_global_round: int) -> None: acc_train = 0.0 for iteration in range(self.local_step): self.model.train() - if hasattr(self.args, "batch_size") and self.args.batch_size > 0: + if use_mini_batch: # print(f"Training with batch size {self.args.batch_size}") loader = NeighborLoader( data, num_neighbors=[-1] * self.args.num_layers, batch_size=self.args.batch_size, input_nodes=self.idx_train, - shuffle=False, + shuffle=True, num_workers=0, ) - batch_iter = iter(loader) - batch = next(batch_iter, None) - while batch is not None: + batch = next(iter(loader), None) + if batch is None: + loss_train, acc_train = 0.0, 0.0 + else: batch_feature_aggregation = batch.x batch_adj_matrix = batch.edge_index + seed_node_count = int(batch.batch_size) + seed_node_index = torch.arange( + seed_node_count, device=batch_feature_aggregation.device + ) + seed_labels = batch.y[:seed_node_count].to( + batch_feature_aggregation.device + ) - # print(f"Batch Feature Aggregation (Node Features): {batch_feature_aggregation.size()}") - # print(f"Batch Adjacency Matrix (Edge Index): {batch_adj_matrix}") - # print(f"Training Labels (Filtered by train_mask): {batch.y[batch.train_mask]}") - # print(f"Train Mask: {batch.train_mask}") + # NeighborLoader puts the sampled seed nodes first. Use + # only those nodes for supervised loss; remaining nodes + # provide GCN message-passing context. loss_train, acc_train = train( iteration, self.model, self.optimizer, batch_feature_aggregation, batch_adj_matrix, - batch.y[batch.train_mask], - batch.train_mask, + seed_labels, + seed_node_index, ) # print(f"acc_train: {acc_train}") - - self.train_losses.append(loss_train) - self.train_accs.append(acc_train) - batch = next(batch_iter, None) else: # print("Training with full batch") # print(f"feature_aggregation size: {self.feature_aggregation.size()}") @@ -916,7 +944,52 @@ def train(self, current_global_round: int) -> None: self.train_losses.append(loss_train) self.train_accs.append(acc_train) # print(f"acc_train: {acc_train}") - self.local_test() + + def _local_eval( + self, + labels: torch.Tensor, + indexes: torch.Tensor, + losses: list, + accuracies: list, + ) -> list: + if self.model is None or self.feature_aggregation is None: + return [0.0, 0.0] + if ( + labels is None + or indexes is None + or labels.numel() == 0 + or indexes.numel() == 0 + ): + return [0.0, 0.0] + + # Ensure everything is on the trainer's device (model may have been + # moved to CPU during aggregation). + self.model = self.model.to(self.device) + feats = self.feature_aggregation.to(self.device) + adj = self.adj.to(self.device) + labels = labels.to(self.device) + indexes = indexes.to(self.device) + local_loss, local_acc = test(self.model, feats, adj, labels, indexes) + losses.append(local_loss) + accuracies.append(local_acc) + return [local_loss, local_acc] + + def local_val(self) -> list: + """ + Evaluates the model on the local validation dataset. + + Returns + ------- + (list) : list + A list containing the validation loss and accuracy + [local_val_loss, local_val_acc]. + """ + return self._local_eval( + self.val_labels, + self.idx_val, + self.val_losses, + self.val_accs, + ) def local_test(self) -> list: """ @@ -927,26 +1000,12 @@ def local_test(self) -> list: (list) : list A list containing the test loss and accuracy [local_test_loss, local_test_acc]. """ - if self.model is None or self.feature_aggregation is None: - return [0.0, 0.0] - - # Ensure everything is on the trainer's device (model may have been - # moved to CPU during aggregation). - self.model = self.model.to(self.device) - feats = self.feature_aggregation.to(self.device) - adj = self.adj.to(self.device) - test_labels = self.test_labels.to(self.device) - idx_test = self.idx_test.to(self.device) - local_test_loss, local_test_acc = test( - self.model, - feats, - adj, - test_labels, - idx_test, + return self._local_eval( + self.test_labels, + self.idx_test, + self.test_losses, + self.test_accs, ) - self.test_losses.append(local_test_loss) - self.test_accs.append(local_test_acc) - return [local_test_loss, local_test_acc] def get_params(self) -> tuple: """ diff --git a/fedgraph/utils_nc.py b/fedgraph/utils_nc.py index eba0adc..d266e84 100644 --- a/fedgraph/utils_nc.py +++ b/fedgraph/utils_nc.py @@ -264,6 +264,7 @@ def get_in_comm_indexes( L_hop: int, idx_train: torch.Tensor, idx_test: torch.Tensor, + idx_val: torch.Tensor = None, ) -> tuple: """ Extract and preprocess data indices and edge information. It determines the nodes that each client @@ -293,6 +294,9 @@ def get_in_comm_indexes( A list of node indices for each client, representing nodes involved in communication. in_com_train_node_indexes : list A list of tensors, where each tensor contains the indices of training data points available to each client. + in_com_val_node_indexes : list + A list of tensors, where each tensor contains the indices of validation + data points available to each client. in_com_test_node_indexes : list A list of tensors, where each tensor contains the indices of test data points available to each client. edge_indexes_clients : list @@ -308,6 +312,7 @@ def get_in_comm_indexes( communicate_node_indexes = [] in_com_train_node_indexes = [] + in_com_val_node_indexes = [] edge_indexes_clients = [] for i in range(num_clients): @@ -362,12 +367,26 @@ def get_in_comm_indexes( torch.searchsorted(communicate_node_indexes[i], inter).clone() ) # local id in block matrix + if idx_val is not None: + inter = intersect1d(split_node_indexes[i], idx_val) + in_com_val_node_indexes.append( + torch.searchsorted(communicate_node_indexes[i], inter).clone() + ) + in_com_test_node_indexes = [] for i in range(num_clients): inter = intersect1d(split_node_indexes[i], idx_test) in_com_test_node_indexes.append( torch.searchsorted(communicate_node_indexes[i], inter).clone() ) + if idx_val is not None: + return ( + communicate_node_indexes, + in_com_train_node_indexes, + in_com_val_node_indexes, + in_com_test_node_indexes, + edge_indexes_clients, + ) return ( communicate_node_indexes, in_com_train_node_indexes, @@ -510,9 +529,11 @@ def save_trainer_data_to_hugging_face( communicate_node_global_index, global_edge_index_client, train_labels, + val_labels, test_labels, features, in_com_train_node_local_indexes, + in_com_val_node_local_indexes, in_com_test_node_local_indexes, global_node_num, class_num, @@ -546,9 +567,11 @@ def save_tensor_to_hf(tensor, file_name): save_tensor_to_hf(communicate_node_global_index, "communicate_node_index.pt") save_tensor_to_hf(global_edge_index_client, "adj.pt") save_tensor_to_hf(train_labels, "train_labels.pt") + save_tensor_to_hf(val_labels, "val_labels.pt") save_tensor_to_hf(test_labels, "test_labels.pt") save_tensor_to_hf(features, "features.pt") save_tensor_to_hf(in_com_train_node_local_indexes, "idx_train.pt") + save_tensor_to_hf(in_com_val_node_local_indexes, "idx_val.pt") save_tensor_to_hf(in_com_test_node_local_indexes, "idx_test.pt") save_tensor_to_hf(torch.tensor(global_node_num), "global_node_num.pt") save_tensor_to_hf(torch.tensor(class_num), "class_num.pt") @@ -563,6 +586,7 @@ def save_all_trainers_data( labels, features, in_com_train_node_local_indexes, + in_com_val_node_local_indexes, in_com_test_node_local_indexes, n_trainer, class_num, @@ -578,11 +602,15 @@ def save_all_trainers_data( train_labels=labels[communicate_node_global_indexes[i]][ in_com_train_node_local_indexes[i] ], + val_labels=labels[communicate_node_global_indexes[i]][ + in_com_val_node_local_indexes[i] + ], test_labels=labels[communicate_node_global_indexes[i]][ in_com_test_node_local_indexes[i] ], features=features[split_node_indexes[i]], in_com_train_node_local_indexes=in_com_train_node_local_indexes[i], + in_com_val_node_local_indexes=in_com_val_node_local_indexes[i], in_com_test_node_local_indexes=in_com_test_node_local_indexes[i], global_node_num=global_node_num, class_num=class_num, diff --git a/tests/integration/test_fedgraph_integration.py b/tests/integration/test_fedgraph_integration.py index 75dbbf6..9949e98 100644 --- a/tests/integration/test_fedgraph_integration.py +++ b/tests/integration/test_fedgraph_integration.py @@ -31,6 +31,7 @@ def setup_method(self): self.labels = torch.randint(0, self.num_classes, (self.num_nodes,)) self.edge_index = torch.randint(0, self.num_nodes, (2, 200)) self.idx_train = torch.arange(0, 70) + self.idx_val = torch.arange(70, 85) self.idx_test = torch.arange(85, 100) # Create node splits for trainers @@ -57,12 +58,18 @@ def test_nc_data_processing_pipeline( from fedgraph.data_process import data_loader_NC # Setup mocks + mock_adj = Mock() + mock_adj.coo.return_value = ( + self.edge_index[0], + self.edge_index[1], + torch.ones(self.edge_index.size(1)), + ) mock_load_data.return_value = ( self.features, - Mock(), + mock_adj, self.labels, self.idx_train, - Mock(), + self.idx_val, self.idx_test, ) mock_partition.return_value = [ @@ -71,6 +78,7 @@ def test_nc_data_processing_pipeline( mock_get_indexes.return_value = ( {i: torch.arange(i * 20, (i + 1) * 20) for i in range(self.num_trainers)}, {i: torch.arange(0, 10) for i in range(self.num_trainers)}, + {i: torch.arange(15, 18) for i in range(self.num_trainers)}, {i: torch.arange(10, 15) for i in range(self.num_trainers)}, {i: self.edge_index for i in range(self.num_trainers)}, ) @@ -87,7 +95,7 @@ def test_nc_data_processing_pipeline( # Test the pipeline result = data_loader_NC(args) - assert len(result) == 11 # Expected return tuple length + assert len(result) == 13 # Expected return tuple length assert mock_load_data.called assert mock_partition.called assert mock_get_indexes.called diff --git a/tests/unit/test_data_process.py b/tests/unit/test_data_process.py index ec538eb..6c7a950 100644 --- a/tests/unit/test_data_process.py +++ b/tests/unit/test_data_process.py @@ -160,6 +160,7 @@ def test_load_cora_dataset( with patch("scipy.sparse.vstack") as mock_vstack, patch( "networkx.adjacency_matrix" ) as mock_adj, patch("networkx.from_dict_of_lists") as mock_from_dict: + mock_features = Mock() mock_features.toarray.return_value = np.random.random((2708, 50)) mock_vstack.return_value = mock_features @@ -355,6 +356,7 @@ def test_data_loader_NC_complete_flow( { i: torch.arange(0, 5) for i in range(3) }, # in_com_train_node_local_indexes + {i: torch.arange(8, 9) for i in range(3)}, # in_com_val_node_local_indexes {i: torch.arange(5, 8) for i in range(3)}, # in_com_test_node_local_indexes { i: torch.stack([torch.arange(0, 10), torch.arange(1, 11)]) @@ -364,17 +366,19 @@ def test_data_loader_NC_complete_flow( result = data_loader_NC(args) - assert len(result) == 11 + assert len(result) == 13 ( edge_index, returned_features, returned_labels, returned_idx_train, + returned_idx_val, returned_idx_test, class_num, split_node_indexes, communicate_node_global_indexes, in_com_train_node_local_indexes, + in_com_val_node_local_indexes, in_com_test_node_local_indexes, global_edge_indexes_clients, ) = result @@ -383,6 +387,7 @@ def test_data_loader_NC_complete_flow( assert torch.equal(returned_features, features) assert torch.equal(returned_labels, labels) assert torch.equal(returned_idx_train, idx_train) + assert torch.equal(returned_idx_val, idx_val) assert torch.equal(returned_idx_test, idx_test) assert class_num == num_classes assert len(split_node_indexes) == args.n_trainer @@ -460,6 +465,7 @@ def test_complete_NC_workflow_mock(self): ) as mock_partition, patch( "fedgraph.data_process.get_in_comm_indexes" ) as mock_indexes: + # Setup test data features = torch.randn(100, 50) adj = torch_sparse.tensor.SparseTensor.from_dense(torch.eye(100)) @@ -477,7 +483,7 @@ def test_complete_NC_workflow_mock(self): idx_test, ) mock_partition.return_value = [list(range(0, 50)), list(range(50, 100))] - mock_indexes.return_value = ({}, {}, {}, {}) + mock_indexes.return_value = ({}, {}, {}, {}, {}) # Create args args = Mock() @@ -493,7 +499,7 @@ def test_complete_NC_workflow_mock(self): result = data_loader(args) assert result is not None - assert len(result) == 11 + assert len(result) == 13 mock_load.assert_called_once() mock_partition.assert_called_once() mock_indexes.assert_called_once() diff --git a/tests/unit/test_federated_methods.py b/tests/unit/test_federated_methods.py index c8239f7..5f12eef 100644 --- a/tests/unit/test_federated_methods.py +++ b/tests/unit/test_federated_methods.py @@ -8,6 +8,8 @@ from fedgraph.federated_methods import ( _resolve_nc_class_num, _resolve_nc_global_node_num, + _unpack_nc_data, + _weighted_nc_metric, run_fedgraph, run_fedgraph_enhanced, run_GC, @@ -83,6 +85,59 @@ def test_rejects_inconsistent_huggingface_global_node_num_metadata(self): _resolve_nc_global_node_num(True, trainer_information) +class TestNCMetricHelpers: + def test_unpack_nc_data_with_validation_fields(self): + data = ( + torch.randn(2, 4), + torch.randn(10, 5), + torch.arange(10), + torch.arange(0, 6), + torch.arange(6, 8), + torch.arange(8, 10), + 3, + [torch.arange(0, 5), torch.arange(5, 10)], + {0: torch.arange(0, 5), 1: torch.arange(5, 10)}, + {0: torch.arange(0, 3), 1: torch.arange(0, 3)}, + {0: torch.arange(3, 4), 1: torch.arange(3, 4)}, + {0: torch.arange(4, 5), 1: torch.arange(4, 5)}, + {0: torch.randn(2, 4), 1: torch.randn(2, 4)}, + ) + + nc_data = _unpack_nc_data(data) + + assert torch.equal(nc_data["idx_val"], data[4]) + assert nc_data["in_com_val_node_local_indexes"] == data[10] + + def test_unpack_nc_data_legacy_tuple_creates_empty_validation_fields(self): + data = ( + torch.randn(2, 4), + torch.randn(10, 5), + torch.arange(10), + torch.arange(0, 6), + torch.arange(8, 10), + 3, + [torch.arange(0, 5), torch.arange(5, 10)], + {0: torch.arange(0, 5), 1: torch.arange(5, 10)}, + {0: torch.arange(0, 3), 1: torch.arange(0, 3)}, + {0: torch.arange(4, 5), 1: torch.arange(4, 5)}, + {0: torch.randn(2, 4), 1: torch.randn(2, 4)}, + ) + + nc_data = _unpack_nc_data(data) + + assert nc_data["idx_val"].numel() == 0 + assert all( + val_indexes.numel() == 0 + for val_indexes in nc_data["in_com_val_node_local_indexes"].values() + ) + + def test_weighted_nc_metric_handles_empty_validation_weights(self): + results = np.array([[1.0, 0.5], [2.0, 0.75]]) + + assert _weighted_nc_metric(results, [0, 0], 1) == 0.0 + assert _weighted_nc_metric(results, [1, 3], 1) == pytest.approx(0.6875) + + class TestRunFedgraph: """Test run_fedgraph main orchestration function.""" @@ -244,6 +299,7 @@ def setup_method(self): self.args.dataset = "cora" self.args.seed = 42 self.args.use_ray = True + self.args.use_cluster = False self.args.he = False self.args.dp = False @@ -283,6 +339,7 @@ def test_run_nc_basic_setup(self, mock_monitor, mock_ray): with patch("fedgraph.federated_methods.Server") as mock_server_class, patch( "fedgraph.federated_methods.torch.manual_seed" ): + mock_server = Mock() mock_server_class.return_value = mock_server @@ -303,14 +360,15 @@ def test_run_nc_basic_setup(self, mock_monitor, mock_ray): mock_monitor.assert_called_once() @patch("fedgraph.federated_methods.ray") - def test_run_nc_without_ray(self, mock_ray): - """Test run_NC without Ray distributed computing.""" + def test_run_nc_initializes_ray_for_execution(self, mock_ray): + """Test run_NC initializes Ray for the current actor-based workflow.""" self.args.use_ray = False mock_ray.init = Mock() with patch("fedgraph.federated_methods.Server") as mock_server_class, patch( "fedgraph.federated_methods.Trainer_General" ) as mock_trainer_class, patch("fedgraph.federated_methods.torch.manual_seed"): + mock_server = Mock() mock_server_class.return_value = mock_server mock_server.trainers = [] @@ -324,7 +382,7 @@ def test_run_nc_without_ray(self, mock_ray): # Expected to fail due to complex flow, but verify no Ray init pass - mock_ray.init.assert_not_called() + mock_ray.init.assert_called() class TestRunGC: @@ -400,6 +458,7 @@ def test_run_gc_gcfl(self, mock_run_gcfl): ) as mock_setup_server, patch( "fedgraph.federated_methods.setup_trainers" ) as mock_setup_trainers: + mock_setup_server.return_value = Mock() mock_setup_trainers.return_value = [Mock(), Mock()] @@ -435,6 +494,7 @@ def test_run_lp_basic_setup( with patch("fedgraph.federated_methods.Server_LP") as mock_server_class, patch( "fedgraph.federated_methods.Monitor" ) as mock_monitor, patch("fedgraph.federated_methods.ray"): + mock_server = Mock() mock_server_class.return_value = mock_server mock_monitor_instance = Mock() diff --git a/tests/unit/test_trainer_class.py b/tests/unit/test_trainer_class.py index 00cc451..9fc03d4 100644 --- a/tests/unit/test_trainer_class.py +++ b/tests/unit/test_trainer_class.py @@ -35,9 +35,11 @@ def test_load_trainer_data_success( torch.randn(50), # communicate_node_global_index torch.randn(2, 200), # global_edge_index_client torch.randn(80), # train_labels + torch.randn(10), # val_labels torch.randn(20), # test_labels torch.randn(100, 10), # features torch.randn(80), # in_com_train_node_local_indexes + torch.randn(10), # in_com_val_node_local_indexes torch.randn(20), # in_com_test_node_local_indexes torch.tensor(100), # global_node_num torch.tensor(3), # class_num @@ -52,11 +54,11 @@ def test_load_trainer_data_success( result = load_trainer_data_from_hugging_face(trainer_id=0, args=args) - assert len(result) == 10 + assert len(result) == 12 assert all(isinstance(tensor, torch.Tensor) for tensor in result) # Verify calls - assert mock_hf_download.call_count == 10 + assert mock_hf_download.call_count == 12 expected_repo = "FedGraph/fedgraph_cora_5trainer_2hop_iid_beta_0.5_trainer_id_0" mock_hf_download.assert_any_call( repo_id=expected_repo, repo_type="dataset", filename="local_node_index.pt" @@ -71,7 +73,7 @@ def test_load_existing_repo_without_global_metadata( mock_file = Mock() mock_file.read.return_value = b"test_tensor_data" mock_open.return_value.__enter__.return_value = mock_file - mock_torch_load.side_effect = [torch.tensor([i]) for i in range(8)] + mock_torch_load.side_effect = [torch.tensor([i]) for i in range(10)] def download_side_effect(*, filename, **kwargs): if filename in {"global_node_num.pt", "class_num.pt"}: @@ -84,7 +86,7 @@ def download_side_effect(*, filename, **kwargs): with pytest.warns(UserWarning, match="falling back to inference"): result = load_trainer_data_from_hugging_face(trainer_id=0, args=args) - assert len(result) == 10 + assert len(result) == 12 assert result[-2:] == (None, None) @@ -156,9 +158,11 @@ def test_trainer_init_without_data(self, mock_load_data): self.communicate_node_index, self.adj, self.train_labels, + torch.empty(0, dtype=self.train_labels.dtype), self.test_labels, self.features, self.idx_train, + torch.empty(0, dtype=self.idx_train.dtype), self.idx_test, torch.tensor(100), torch.tensor(3), @@ -434,11 +438,62 @@ def test_train_method(self, mock_train_func, mock_test_func): trainer.train(current_global_round=1) assert mock_train_func.call_count == trainer.local_step - assert mock_test_func.call_count == trainer.local_step + assert mock_test_func.call_count == 0 assert len(trainer.train_losses) == trainer.local_step assert len(trainer.train_accs) == trainer.local_step - assert len(trainer.test_losses) == trainer.local_step - assert len(trainer.test_accs) == trainer.local_step + assert len(trainer.test_losses) == 0 + assert len(trainer.test_accs) == 0 + + @patch("fedgraph.trainer_class.NeighborLoader") + @patch("fedgraph.trainer_class.train") + def test_train_method_mini_batch_one_update_per_local_step( + self, mock_train_func, mock_neighbor_loader + ): + """Mini-batch NC training should use one seed batch per local step.""" + trainer = Trainer_General( + rank=self.rank, + args_hidden=self.args_hidden, + device=self.device, + args=self.args, + local_node_index=self.local_node_index, + communicate_node_index=self.communicate_node_index, + adj=self.adj, + train_labels=self.train_labels, + test_labels=self.test_labels, + features=self.features, + idx_train=self.idx_train, + idx_test=self.idx_test, + ) + + trainer.model = Mock() + trainer.optimizer = Mock() + trainer.class_num = 7 + self.args.batch_size = 2 + + batch = Mock() + batch.batch_size = self.args.batch_size + batch.x = torch.randn(5, self.features.shape[1]) + batch.edge_index = torch.tensor([[0, 1, 2], [1, 2, 3]]) + batch.y = torch.tensor([0, 1, 2, -1, -1]) + mock_neighbor_loader.return_value = [batch] + mock_train_func.return_value = (0.5, 0.85) + + trainer.train(current_global_round=1) + + assert mock_neighbor_loader.call_count == trainer.local_step + assert mock_train_func.call_count == trainer.local_step + assert len(trainer.train_losses) == trainer.local_step + assert len(trainer.train_accs) == trainer.local_step + + for call in mock_neighbor_loader.call_args_list: + assert call.kwargs["batch_size"] == self.args.batch_size + assert call.kwargs["shuffle"] is True + assert torch.equal(call.kwargs["input_nodes"], trainer.idx_train) + + expected_seed_index = torch.arange(self.args.batch_size) + for call in mock_train_func.call_args_list: + assert torch.equal(call.args[5], batch.y[: self.args.batch_size]) + assert torch.equal(call.args[6].cpu(), expected_seed_index) @patch("fedgraph.trainer_class.test") def test_local_test(self, mock_test_func): diff --git a/tests/unit/test_utils_nc.py b/tests/unit/test_utils_nc.py index 634f843..079e9e4 100644 --- a/tests/unit/test_utils_nc.py +++ b/tests/unit/test_utils_nc.py @@ -131,9 +131,11 @@ def test_uploads_global_metadata(self, mock_hf_api, mock_get_token): communicate_node_global_index=torch.tensor([0, 1, 2]), global_edge_index_client=torch.tensor([[0, 1], [1, 2]]), train_labels=torch.tensor([0]), + val_labels=torch.tensor([2]), test_labels=torch.tensor([1]), features=torch.randn(2, 4), in_com_train_node_local_indexes=torch.tensor([0]), + in_com_val_node_local_indexes=torch.tensor([2]), in_com_test_node_local_indexes=torch.tensor([1]), global_node_num=3, class_num=2, @@ -145,6 +147,8 @@ def test_uploads_global_metadata(self, mock_hf_api, mock_get_token): } assert "global_node_num.pt" in uploaded_files assert "class_num.pt" in uploaded_files + assert "val_labels.pt" in uploaded_files + assert "idx_val.pt" in uploaded_files @patch("fedgraph.utils_nc.save_trainer_data_to_hugging_face") def test_bulk_save_passes_consistent_global_metadata(self, mock_save_trainer): @@ -154,8 +158,8 @@ def test_bulk_save_passes_consistent_global_metadata(self, mock_save_trainer): save_all_trainers_data( split_node_indexes=[torch.tensor([0, 1]), torch.tensor([2, 3])], communicate_node_global_indexes=[ - torch.tensor([0, 1]), - torch.tensor([2, 3]), + torch.tensor([0, 1, 2]), + torch.tensor([1, 2, 3]), ], global_edge_indexes_clients=[ torch.tensor([[0], [1]]), @@ -164,6 +168,7 @@ def test_bulk_save_passes_consistent_global_metadata(self, mock_save_trainer): labels=labels, features=features, in_com_train_node_local_indexes=[torch.tensor([0]), torch.tensor([0])], + in_com_val_node_local_indexes=[torch.tensor([2]), torch.tensor([2])], in_com_test_node_local_indexes=[torch.tensor([1]), torch.tensor([1])], n_trainer=2, class_num=3, @@ -174,6 +179,10 @@ def test_bulk_save_passes_consistent_global_metadata(self, mock_save_trainer): for call in mock_save_trainer.call_args_list: assert call.kwargs["global_node_num"] == 4 assert call.kwargs["class_num"] == 3 + assert not torch.equal( + call.kwargs["in_com_val_node_local_indexes"], + call.kwargs["in_com_test_node_local_indexes"], + ) class TestLabelDirichletPartition: @@ -320,6 +329,40 @@ def test_get_in_comm_indexes_basic(self): assert len(communicate_node_global_indexes) == n_trainers assert len(global_edge_indexes_clients) == n_trainers + def test_get_in_comm_indexes_with_validation_indexes(self): + """Test communication index generation includes validation nodes.""" + edge_index = torch.tensor([[0, 1, 2, 3, 4], [1, 2, 3, 4, 0]]) + split_node_indexes = [ + torch.tensor([0, 1]), + torch.tensor([2, 3]), + torch.tensor([4]), + ] + + result = get_in_comm_indexes( + edge_index, + split_node_indexes, + 3, + 2, + torch.tensor([0, 2]), + torch.tensor([1, 3]), + idx_val=torch.tensor([4]), + ) + + assert len(result) == 5 + ( + communicate_node_global_indexes, + in_com_train_node_local_indexes, + in_com_val_node_local_indexes, + in_com_test_node_local_indexes, + global_edge_indexes_clients, + ) = result + + assert len(communicate_node_global_indexes) == 3 + assert len(in_com_train_node_local_indexes) == 3 + assert len(in_com_val_node_local_indexes) == 3 + assert len(in_com_test_node_local_indexes) == 3 + assert len(global_edge_indexes_clients) == 3 + def test_get_in_comm_indexes_rejects_unsupported_one_hop(self): """Test that the old ambiguous 1-hop NC mode is rejected.""" edge_index = torch.tensor([[0, 1], [1, 2]])