diff --git a/tests/test_risk_compare_batch_iterative.py b/tests/test_risk_compare_batch_iterative.py index c6ad2701..4de3fab3 100644 --- a/tests/test_risk_compare_batch_iterative.py +++ b/tests/test_risk_compare_batch_iterative.py @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. - +import numbers import unittest import datetime import pytz @@ -97,12 +97,16 @@ class RiskCompareIterativeToBatch(unittest.TestCase): self.assertEqual( risk_metrics_original.treasury_period_return, risk_metrics_refactor.treasury_period_return) - self.assertEqual( + np.testing.assert_allclose( risk_metrics_original.benchmark_returns, - risk_metrics_refactor.benchmark_returns) - self.assertEqual( + risk_metrics_refactor.benchmark_returns, + rtol=0.001 + ) + np.testing.assert_allclose( risk_metrics_original.algorithm_returns, - risk_metrics_refactor.algorithm_returns) + risk_metrics_refactor.algorithm_returns, + rtol=0.001 + ) risk_original_dict = risk_metrics_original.to_dict() risk_refactor_dict = risk_metrics_refactor.to_dict() self.assertEqual(set(risk_original_dict.keys()), @@ -122,12 +126,24 @@ class RiskCompareIterativeToBatch(unittest.TestCase): truth=risk_original_dict[measure], returned=risk_refactor_dict[measure])) else: - np.testing.assert_equal( - risk_original_dict[measure], - risk_refactor_dict[measure], - err_msg_format.format( - iter=i, - measure=measure, - truth=risk_original_dict[measure], - returned=risk_refactor_dict[measure]) - ) + if isinstance(risk_original_dict[measure], numbers.Real): + np.testing.assert_allclose( + risk_original_dict[measure], + risk_refactor_dict[measure], + rtol=0.001, + err_msg=err_msg_format.format( + iter=i, + measure=measure, + truth=risk_original_dict[measure], + returned=risk_refactor_dict[measure]) + ) + else: + np.testing.assert_equal( + risk_original_dict[measure], + risk_refactor_dict[measure], + err_msg=err_msg_format.format( + iter=i, + measure=measure, + truth=risk_original_dict[measure], + returned=risk_refactor_dict[measure]) + )