fixed tests

This commit is contained in:
Aakash Kaushik
2020-09-27 23:15:07 +05:30
parent 206409fb1a
commit 790d7122b8
+13 -13
View File
@@ -55,7 +55,7 @@ TEST_CASE("HuberLossTest", "[LossFunctionsTest]")
input = arma::mat("17.45 12.91 13.63 29.01 7.12 15.47 31.52 31.97");
target = arma::mat("16.52 13.11 13.67 29.51 24.31 15.03 30.72 34.07");
double loss = module.Forward(input, target);
REQUIRE(loss == Approx(2.410631).epsilon(1e-10));
REQUIRE(loss == Approx(2.410631).epsilon(1e-6));
// Test the backward function.
module.Backward(input, target, output);
@@ -64,7 +64,7 @@ TEST_CASE("HuberLossTest", "[LossFunctionsTest]")
// [0.1162 -0.0250 -0.0050 -0.0625 -0.1250 0.0550 0.1000 -0.1250]
// Sum of Expected Output = -0.07125.
double expectedOutputSum = arma::accu(output);
REQUIRE(expectedOutputSum == Approx(-0.07125).epsilon(1e-10));
REQUIRE(expectedOutputSum == Approx(-0.07125).epsilon(1e-5));
REQUIRE(output.n_rows == input.n_rows);
REQUIRE(output.n_cols == input.n_cols);
@@ -96,10 +96,10 @@ TEST_CASE("PoissonNLLLossTest", "[LossFunctionsTest]")
double loss2 = module2.Forward(input, target);
double loss3 = module3.Forward(input, target);
double loss4 = module4.Forward(input4, target4);
REQUIRE(loss1 == Approx(4.8986).epsilon(1e-8));
REQUIRE(loss2 == Approx(45.4139).epsilon(1e-8));
REQUIRE(loss3 == Approx(5.6767).epsilon(1e-8));
REQUIRE(loss4 == Approx(3.742157).epsilon(1e-8));
REQUIRE(loss1 == Approx(4.8986).epsilon(1e-4));
REQUIRE(loss2 == Approx(45.4139).epsilon(1e-4));
REQUIRE(loss3 == Approx(5.6767).epsilon(1e-4));
REQUIRE(loss4 == Approx(3.742157).epsilon(1e-4));
// Test the Backward function.
module1.Backward(input, target, output1);
@@ -130,10 +130,10 @@ TEST_CASE("PoissonNLLLossTest", "[LossFunctionsTest]")
for (size_t i = 0; i < expOutput1.n_elem; ++i)
{
REQUIRE(output1[i] == Approx(expOutput1[i]).epsilon(1e-8));
REQUIRE(output2[i] == Approx(expOutput2[i]).epsilon(1e-8));
REQUIRE(output3[i] == Approx(expOutput3[i]).epsilon(1e-8));
REQUIRE(output4[i] == Approx(expOutput4[i]).epsilon(1e-8));
REQUIRE(output1[i] == Approx(expOutput1[i]).epsilon(1e-4));
REQUIRE(output2[i] == Approx(expOutput2[i]).epsilon(1e-4));
REQUIRE(output3[i] == Approx(expOutput3[i]).epsilon(1e-4));
REQUIRE(output4[i] == Approx(expOutput4[i]).epsilon(1e-4));
}
}
@@ -205,7 +205,7 @@ TEST_CASE("KLDivergenceMeanTest", "[LossFunctionsTest]")
// Test the Backward function.
module.Backward(input, target, output);
REQUIRE(arma::as_scalar(output) == Approx(-0.1).epsilon(1e-10));
REQUIRE(arma::as_scalar(output) == Approx(-0.1).epsilon(1e-5));
}
/**
@@ -222,11 +222,11 @@ TEST_CASE("KLDivergenceNoMeanTest", "[LossFunctionsTest]")
target = arma::exp(arma::mat("2 1 1 1 1 1 1 1 1 1"));
loss = module.Forward(input, target);
REQUIRE(loss == Approx(-11).epsilon(1e-10));
REQUIRE(loss == Approx(-11).epsilon(1e-5));
// Test the Backward function.
module.Backward(input, target, output);
REQUIRE(arma::as_scalar(output) == Approx(-1).epsilon(1e-10));
REQUIRE(arma::as_scalar(output) == Approx(-1).epsilon(1e-5));
}
/*