From 7c92b4a731e61332d813b54bfe9d6808494642e4 Mon Sep 17 00:00:00 2001 From: Naoto Mizuno Date: Thu, 12 Oct 2023 17:11:08 +0900 Subject: [PATCH] Add test for _one_side_trunc_norm_sampling --- python_tests/preferential/samplers/test_gp.py | 29 +++++++++++++++++++ 1 file changed, 29 insertions(+) create mode 100644 python_tests/preferential/samplers/test_gp.py diff --git a/python_tests/preferential/samplers/test_gp.py b/python_tests/preferential/samplers/test_gp.py new file mode 100644 index 00000000..a249f4b2 --- /dev/null +++ b/python_tests/preferential/samplers/test_gp.py @@ -0,0 +1,29 @@ +import sys +from unittest.mock import patch + +import numpy as np +import pytest +import torch + + +if sys.version_info >= (3, 8): + from optuna_dashboard.preferential.samplers.gp import _one_side_trunc_norm_sampling +else: + pytest.skip("BoTorch dropped Python3.7 support", allow_module_level=True) + + +def test_one_side_trunc_norm_sampling() -> None: + for lower in np.linspace(-10, 10, 100): + assert _one_side_trunc_norm_sampling(torch.tensor([lower], dtype=torch.float64)) >= lower + + with patch.object(torch, "rand", return_value=torch.tensor([0.4], dtype=torch.float64)): + sampled_value = _one_side_trunc_norm_sampling(torch.tensor([0.1], dtype=torch.float64)) + assert np.allclose(sampled_value.numpy(), 0.899967154837563) + + with patch.object(torch, "rand", return_value=torch.tensor([0.8], dtype=torch.float64)): + sampled_value = _one_side_trunc_norm_sampling(torch.tensor([-2.3], dtype=torch.float64)) + assert np.allclose(sampled_value.numpy(), -0.8113606739551955) + + with patch.object(torch, "rand", return_value=torch.tensor([0.1], dtype=torch.float64)): + sampled_value = _one_side_trunc_norm_sampling(torch.tensor([5], dtype=torch.float64)) + assert np.allclose(sampled_value.numpy(), 5.426934003050024)