diff --git a/tests/test_optim.py b/tests/test_optim.py index 29736311d..eda3a094c 100644 --- a/tests/test_optim.py +++ b/tests/test_optim.py @@ -548,8 +548,15 @@ def test_optimizer8bit(dim1, dim2, gtype, optim_name, device): @pytest.mark.parametrize("dim2", [4096], ids=id_formatter("dim2")) @pytest.mark.parametrize("gtype", [torch.float32, torch.bfloat16, torch.float16], ids=describe_dtype) @pytest.mark.parametrize("optim_name", optimizer_names_benchmark, ids=id_formatter("opt")) +@pytest.mark.parametrize("device", get_available_devices(), ids=id_formatter("device")) @pytest.mark.benchmark def test_benchmark_blockwise(dim1, dim2, gtype, optim_name, device): + if device == "cpu" and optim_name.startswith("paged_"): + pytest.skip("Paged optimizers are not meaningful on CPU") + + if optim_name.startswith("paged_") and sys.platform == "win32": + pytest.skip("Paged optimizers can have issues on Windows.") + if dim1 == 1 and dim2 == 1: return p1 = torch.randn(dim1, dim2, device=device, dtype=gtype) * 0.1