Merge pull request #2875 from RishabhGarg108/TemplateDataStyleFix
Fixed variable names in #2746
This commit is contained in:
@@ -223,9 +223,9 @@ TEST_CASE("ZeroRatioStratifiedSplitData", "[SplitDataTest]")
|
||||
|
||||
// Set the labels to 5 0s and 10 1s.
|
||||
const Row<size_t> labels = { 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 };
|
||||
const double test_ratio = 0;
|
||||
const double testRatio = 0;
|
||||
|
||||
const auto value = Split(input, labels, test_ratio, false, true);
|
||||
const auto value = Split(input, labels, testRatio, false, true);
|
||||
REQUIRE(std::get<0>(value).n_cols == 15);
|
||||
REQUIRE(std::get<1>(value).n_cols == 0);
|
||||
REQUIRE(std::get<2>(value).n_cols == 15);
|
||||
@@ -242,9 +242,9 @@ TEST_CASE("TotalRatioStratifiedSplitData", "[SplitDataTest]")
|
||||
|
||||
// Set the labels to 5 0s and 10 1s.
|
||||
const Row<size_t> labels = { 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 1, 1 };
|
||||
const double test_ratio = 1;
|
||||
const double testRatio = 1;
|
||||
|
||||
const auto value = Split(input, labels, test_ratio, false, true);
|
||||
const auto value = Split(input, labels, testRatio, false, true);
|
||||
REQUIRE(std::get<0>(value).n_cols == 0);
|
||||
REQUIRE(std::get<1>(value).n_cols == 15);
|
||||
REQUIRE(std::get<2>(value).n_cols == 0);
|
||||
@@ -263,9 +263,9 @@ TEST_CASE("StratifiedSplitDataResultTest", "[SplitDataTest]")
|
||||
const Row<size_t> labels = { 0, 0, 0, 0,
|
||||
1, 1, 1, 1, 1, 1, 1, 1,
|
||||
2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2, 2 };
|
||||
const double test_ratio = 0.25;
|
||||
const double testRatio = 0.25;
|
||||
|
||||
const auto value = Split(input, labels, test_ratio, true, true);
|
||||
const auto value = Split(input, labels, testRatio, true, true);
|
||||
REQUIRE(static_cast<uvec>(find(std::get<2>(value) == 0)).n_rows == 3);
|
||||
REQUIRE(static_cast<uvec>(find(std::get<2>(value) == 1)).n_rows == 6);
|
||||
REQUIRE(static_cast<uvec>(find(std::get<2>(value) == 2)).n_rows == 9);
|
||||
@@ -293,22 +293,22 @@ TEST_CASE("StratifiedSplitLargerDataResultTest", "[SplitDataTest]")
|
||||
input.randu();
|
||||
|
||||
// 256 0s, 128 1s, 64 2s and 32 3s.
|
||||
Row<size_t> zero_label(256);
|
||||
Row<size_t> one_label(128);
|
||||
Row<size_t> two_label(64);
|
||||
Row<size_t> three_label(32);
|
||||
Row<size_t> zeroLabel(256);
|
||||
Row<size_t> oneLabel(128);
|
||||
Row<size_t> twoLabel(64);
|
||||
Row<size_t> threeLabel(32);
|
||||
|
||||
zero_label.fill(0);
|
||||
one_label.fill(1);
|
||||
two_label.fill(2);
|
||||
three_label.fill(3);
|
||||
zeroLabel.fill(0);
|
||||
oneLabel.fill(1);
|
||||
twoLabel.fill(2);
|
||||
threeLabel.fill(3);
|
||||
|
||||
Row<size_t> labels = arma::join_rows(zero_label, one_label);
|
||||
labels = arma::join_rows(labels, two_label);
|
||||
labels = arma::join_rows(labels, three_label);
|
||||
const double test_ratio = 0.3;
|
||||
Row<size_t> labels = arma::join_rows(zeroLabel, oneLabel);
|
||||
labels = arma::join_rows(labels, twoLabel);
|
||||
labels = arma::join_rows(labels, threeLabel);
|
||||
const double testRatio = 0.3;
|
||||
|
||||
const auto value = Split(input, labels, test_ratio, false, true);
|
||||
const auto value = Split(input, labels, testRatio, false, true);
|
||||
REQUIRE(static_cast<uvec>(find(std::get<2>(value) == 0)).n_rows == 180);
|
||||
REQUIRE(static_cast<uvec>(find(std::get<2>(value) == 1)).n_rows == 90);
|
||||
REQUIRE(static_cast<uvec>(find(std::get<2>(value) == 2)).n_rows == 45);
|
||||
@@ -334,9 +334,9 @@ TEST_CASE("StratifiedSplitRunTimeErrorTest", "[SplitDataTest]")
|
||||
input.randu();
|
||||
labels.randu();
|
||||
|
||||
const double test_ratio = 0.3;
|
||||
const double testRatio = 0.3;
|
||||
|
||||
REQUIRE_THROWS_AS(Split(input, labels, test_ratio, false, true),
|
||||
REQUIRE_THROWS_AS(Split(input, labels, testRatio, false, true),
|
||||
std::runtime_error);
|
||||
}
|
||||
|
||||
@@ -380,12 +380,12 @@ TEST_CASE("SplitMatrixLabeledData", "[SplitDataTest]")
|
||||
REQUIRE(std::get<2>(value).n_cols == 8);
|
||||
REQUIRE(std::get<3>(value).n_cols == 2);
|
||||
|
||||
mat input_concat = arma::join_rows(std::get<0>(value), std::get<1>(value));
|
||||
mat labels_concat = arma::join_rows(std::get<2>(value), std::get<3>(value));
|
||||
mat inputConcat = arma::join_rows(std::get<0>(value), std::get<1>(value));
|
||||
mat labelsConcat = arma::join_rows(std::get<2>(value), std::get<3>(value));
|
||||
|
||||
// Order matters here.
|
||||
CheckMatrices(input, input_concat);
|
||||
CheckMatrices(labels, labels_concat);
|
||||
CheckMatrices(input, inputConcat);
|
||||
CheckMatrices(labels, labelsConcat);
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -414,10 +414,10 @@ TEST_CASE("SplitLabeledDataResultField", "[SplitDataTest]")
|
||||
REQUIRE(std::get<2>(value).n_cols == 1); // Train label.
|
||||
REQUIRE(std::get<3>(value).n_cols == 1); // Test label.
|
||||
|
||||
field<mat> input_concat = {std::get<0>(value)(0), std::get<1>(value)(0)};
|
||||
field<vec> label_concat = {std::get<2>(value)(0), std::get<3>(value)(0)};
|
||||
field<mat> inputConcat = {std::get<0>(value)(0), std::get<1>(value)(0)};
|
||||
field<vec> labelConcat = {std::get<2>(value)(0), std::get<3>(value)(0)};
|
||||
|
||||
// Order matters here.
|
||||
CheckFields(input, input_concat);
|
||||
CheckFields(label, label_concat);
|
||||
CheckFields(input, inputConcat);
|
||||
CheckFields(label, labelConcat);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user