merge 23Mar01
This commit is contained in:
@@ -13,25 +13,6 @@
|
||||
#include "../../../plugin/federated/federated_server.h"
|
||||
#include "../../../src/collective/communicator-inl.h"
|
||||
|
||||
inline int GenerateRandomPort(int low, int high) {
|
||||
using namespace std::chrono_literals;
|
||||
// Ensure unique timestamp by introducing a small artificial delay
|
||||
std::this_thread::sleep_for(100ms);
|
||||
auto timestamp = static_cast<uint64_t>(std::chrono::duration_cast<std::chrono::milliseconds>(
|
||||
std::chrono::system_clock::now().time_since_epoch())
|
||||
.count());
|
||||
std::mt19937_64 rng(timestamp);
|
||||
std::uniform_int_distribution<int> dist(low, high);
|
||||
int port = dist(rng);
|
||||
return port;
|
||||
}
|
||||
|
||||
inline std::string GetServerAddress() {
|
||||
int port = GenerateRandomPort(50000, 60000);
|
||||
std::string address = std::string("localhost:") + std::to_string(port);
|
||||
return address;
|
||||
}
|
||||
|
||||
namespace xgboost {
|
||||
|
||||
class ServerForTest {
|
||||
@@ -41,13 +22,14 @@ class ServerForTest {
|
||||
|
||||
public:
|
||||
explicit ServerForTest(std::int32_t world_size) {
|
||||
server_address_ = GetServerAddress();
|
||||
server_thread_.reset(new std::thread([this, world_size] {
|
||||
grpc::ServerBuilder builder;
|
||||
xgboost::federated::FederatedService service{world_size};
|
||||
builder.AddListeningPort(server_address_, grpc::InsecureServerCredentials());
|
||||
int selected_port;
|
||||
builder.AddListeningPort("localhost:0", grpc::InsecureServerCredentials(), &selected_port);
|
||||
builder.RegisterService(&service);
|
||||
server_ = builder.BuildAndStart();
|
||||
server_address_ = std::string("localhost:") + std::to_string(selected_port);
|
||||
server_->Wait();
|
||||
}));
|
||||
}
|
||||
@@ -56,7 +38,14 @@ class ServerForTest {
|
||||
server_->Shutdown();
|
||||
server_thread_->join();
|
||||
}
|
||||
auto Address() const { return server_address_; }
|
||||
|
||||
auto Address() const {
|
||||
using namespace std::chrono_literals;
|
||||
while (server_address_.empty()) {
|
||||
std::this_thread::sleep_for(100ms);
|
||||
}
|
||||
return server_address_;
|
||||
}
|
||||
};
|
||||
|
||||
class BaseFederatedTest : public ::testing::Test {
|
||||
@@ -65,7 +54,7 @@ class BaseFederatedTest : public ::testing::Test {
|
||||
|
||||
void TearDown() override { server_.reset(nullptr); }
|
||||
|
||||
static int const kWorldSize{3};
|
||||
static int constexpr kWorldSize{3};
|
||||
std::unique_ptr<ServerForTest> server_;
|
||||
};
|
||||
|
||||
|
||||
@@ -62,34 +62,24 @@ class FederatedCommunicatorTest : public BaseFederatedTest {
|
||||
};
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, ThrowOnWorldSizeTooSmall) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
auto construct = [server_address]() {
|
||||
FederatedCommunicator comm{0, 0, server_address, "", "", ""};
|
||||
};
|
||||
auto construct = [] { FederatedCommunicator comm{0, 0, "localhost:0", "", "", ""}; };
|
||||
EXPECT_THROW(construct(), dmlc::Error);
|
||||
}
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, ThrowOnRankTooSmall) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
auto construct = [server_address]() {
|
||||
FederatedCommunicator comm{1, -1, server_address, "", "", ""};
|
||||
};
|
||||
auto construct = [] { FederatedCommunicator comm{1, -1, "localhost:0", "", "", ""}; };
|
||||
EXPECT_THROW(construct(), dmlc::Error);
|
||||
}
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, ThrowOnRankTooBig) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
auto construct = [server_address]() {
|
||||
FederatedCommunicator comm{1, 1, server_address, "", "", ""};
|
||||
};
|
||||
auto construct = [] { FederatedCommunicator comm{1, 1, "localhost:0", "", "", ""}; };
|
||||
EXPECT_THROW(construct(), dmlc::Error);
|
||||
}
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, ThrowOnWorldSizeNotInteger) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
auto construct = [server_address]() {
|
||||
auto construct = [] {
|
||||
Json config{JsonObject()};
|
||||
config["federated_server_address"] = server_address;
|
||||
config["federated_server_address"] = std::string("localhost:0");
|
||||
config["federated_world_size"] = std::string("1");
|
||||
config["federated_rank"] = Integer(0);
|
||||
FederatedCommunicator::Create(config);
|
||||
@@ -98,10 +88,9 @@ TEST(FederatedCommunicatorSimpleTest, ThrowOnWorldSizeNotInteger) {
|
||||
}
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, ThrowOnRankNotInteger) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
auto construct = [server_address]() {
|
||||
auto construct = [] {
|
||||
Json config{JsonObject()};
|
||||
config["federated_server_address"] = server_address;
|
||||
config["federated_server_address"] = std::string("localhost:0");
|
||||
config["federated_world_size"] = 1;
|
||||
config["federated_rank"] = std::string("0");
|
||||
FederatedCommunicator::Create(config);
|
||||
@@ -110,15 +99,13 @@ TEST(FederatedCommunicatorSimpleTest, ThrowOnRankNotInteger) {
|
||||
}
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, GetWorldSizeAndRank) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
FederatedCommunicator comm{6, 3, server_address};
|
||||
FederatedCommunicator comm{6, 3, "localhost:0"};
|
||||
EXPECT_EQ(comm.GetWorldSize(), 6);
|
||||
EXPECT_EQ(comm.GetRank(), 3);
|
||||
}
|
||||
|
||||
TEST(FederatedCommunicatorSimpleTest, IsDistributed) {
|
||||
std::string server_address{GetServerAddress()};
|
||||
FederatedCommunicator comm{2, 1, server_address};
|
||||
FederatedCommunicator comm{2, 1, "localhost:0"};
|
||||
EXPECT_TRUE(comm.IsDistributed());
|
||||
}
|
||||
|
||||
|
||||
@@ -70,7 +70,7 @@ void VerifyObjective(size_t rows, size_t cols, float expected_base_score, Json e
|
||||
|
||||
class FederatedLearnerTest : public ::testing::TestWithParam<std::string> {
|
||||
std::unique_ptr<ServerForTest> server_;
|
||||
static int const kWorldSize{3};
|
||||
static int constexpr kWorldSize{3};
|
||||
|
||||
protected:
|
||||
void SetUp() override { server_ = std::make_unique<ServerForTest>(kWorldSize); }
|
||||
|
||||
243
tests/cpp/plugin/test_federated_metrics.cc
Normal file
243
tests/cpp/plugin/test_federated_metrics.cc
Normal file
@@ -0,0 +1,243 @@
|
||||
/*!
|
||||
* Copyright 2023 XGBoost contributors
|
||||
*/
|
||||
#include <gtest/gtest.h>
|
||||
|
||||
#include "../metric/test_auc.h"
|
||||
#include "../metric/test_elementwise_metric.h"
|
||||
#include "../metric/test_multiclass_metric.h"
|
||||
#include "../metric/test_rank_metric.h"
|
||||
#include "../metric/test_survival_metric.h"
|
||||
#include "helpers.h"
|
||||
|
||||
namespace {
|
||||
class FederatedMetricTest : public xgboost::BaseFederatedTest {};
|
||||
} // anonymous namespace
|
||||
|
||||
namespace xgboost {
|
||||
namespace metric {
|
||||
TEST_F(FederatedMetricTest, BinaryAUCRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyBinaryAUC,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, BinaryAUCColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyBinaryAUC,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassAUCRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassAUC,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassAUCColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassAUC,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RankingAUCRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRankingAUC,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RankingAUCColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRankingAUC,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, PRAUCRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyPRAUC, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, PRAUCColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyPRAUC, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassPRAUCRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassPRAUC,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassPRAUCColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassPRAUC,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RankingPRAUCRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRankingPRAUC,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RankingPRAUCColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRankingPRAUC,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RMSERowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRMSE, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RMSEColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRMSE, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RMSLERowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRMSLE, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, RMSLEColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyRMSLE, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MAERowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMAE, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MAEColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMAE, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MAPERowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMAPE, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MAPEColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMAPE, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MPHERowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMPHE, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MPHEColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMPHE, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, LogLossRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyLogLoss, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, LogLossColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyLogLoss, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, ErrorRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyError, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, ErrorColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyError, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, PoissonNegLogLikRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyPoissonNegLogLik,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, PoissonNegLogLikColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyPoissonNegLogLik,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiRMSERowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiRMSE,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiRMSEColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiRMSE,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, QuantileRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyQuantile,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, QuantileColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyQuantile,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassErrorRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassError,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassErrorColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassError,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassLogLossRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassLogLoss,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MultiClassLogLossColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMultiClassLogLoss,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, PrecisionRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyPrecision,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, PrecisionColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyPrecision,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, NDCGRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyNDCG, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, NDCGColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyNDCG, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MAPRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMAP, DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, MAPColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyMAP, DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, NDCGExpGainRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyNDCGExpGain,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, NDCGExpGainColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyNDCGExpGain,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
} // namespace metric
|
||||
} // namespace xgboost
|
||||
|
||||
namespace xgboost {
|
||||
namespace common {
|
||||
TEST_F(FederatedMetricTest, AFTNegLogLikRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyAFTNegLogLik,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, AFTNegLogLikColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyAFTNegLogLik,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, IntervalRegressionAccuracyRowSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyIntervalRegressionAccuracy,
|
||||
DataSplitMode::kRow);
|
||||
}
|
||||
|
||||
TEST_F(FederatedMetricTest, IntervalRegressionAccuracyColumnSplit) {
|
||||
RunWithFederatedCommunicator(kWorldSize, server_->Address(), &VerifyIntervalRegressionAccuracy,
|
||||
DataSplitMode::kCol);
|
||||
}
|
||||
} // namespace common
|
||||
} // namespace xgboost
|
||||
Reference in New Issue
Block a user