mirror of
https://github.com/wiltodelta/remove-ai-watermarks.git
synced 2026-07-31 11:37:22 +02:00
feat(device): support xpu backend (#24)
* feat(device): support xpu backend * Fall back to CPU seed generator when device RNG unsupported (xpu) Some torch-xpu builds have no device-side RNG, so torch.Generator(device="xpu") raises when --seed is used. _make_seed_generator tries the device generator and falls back to a backend-agnostic CPU generator. Adds a fallback unit test. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Victor Kuznetsov <kuznetsov.va@gmail.com> Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
Victor Kuznetsov
parent
1598c499fe
commit
0c7ff1874e
+45
-2
@@ -6,7 +6,7 @@ code paths work correctly on CPU, MPS (macOS), and CUDA (Linux/Windows).
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -27,7 +27,7 @@ class TestDeviceDetection:
|
||||
|
||||
def test_returns_valid_device(self):
|
||||
device = get_device()
|
||||
assert device in ("cpu", "mps", "cuda")
|
||||
assert device in ("cpu", "mps", "cuda", "xpu")
|
||||
|
||||
def test_cpu_fallback_when_no_gpu(self):
|
||||
"""On CI / machines without GPU, should fall back to cpu or mps."""
|
||||
@@ -39,6 +39,49 @@ class TestDeviceDetection:
|
||||
def test_no_torch_returns_cpu(self):
|
||||
assert get_device() == "cpu"
|
||||
|
||||
def test_xpu_selected_when_available(self):
|
||||
"""An XPU-enabled torch (no CUDA) routes to the Intel GPU backend.
|
||||
|
||||
The whole torch module is mocked so the smoke-test ops succeed without
|
||||
any real device; cuda must read False so the cuda branch is skipped.
|
||||
"""
|
||||
fake_torch = MagicMock()
|
||||
fake_torch.cuda.is_available.return_value = False
|
||||
fake_torch.xpu.is_available.return_value = True
|
||||
with patch("remove_ai_watermarks.noai.watermark_remover.torch", fake_torch):
|
||||
assert get_device() == "xpu"
|
||||
fake_torch.tensor.assert_called_with([1.0], device="xpu")
|
||||
|
||||
def test_init_accepts_xpu_and_selects_fp16(self):
|
||||
"""WatermarkRemover accepts device='xpu' and picks fp16 (not fp32)."""
|
||||
if not is_watermark_removal_available():
|
||||
pytest.skip("torch/diffusers not installed")
|
||||
import torch
|
||||
|
||||
from remove_ai_watermarks.noai.watermark_remover import WatermarkRemover
|
||||
|
||||
remover = WatermarkRemover(device="xpu")
|
||||
assert remover.device == "xpu"
|
||||
assert remover.torch_dtype == torch.float16
|
||||
|
||||
def test_seed_generator_falls_back_to_cpu_when_device_rng_unsupported(self):
|
||||
"""A device with no RNG backend (e.g. some torch-xpu builds) falls back
|
||||
to a CPU generator instead of raising when --seed is used."""
|
||||
from remove_ai_watermarks.noai import watermark_remover as wr
|
||||
|
||||
def fake_generator(device="cpu"):
|
||||
if device == "xpu":
|
||||
raise RuntimeError("Device type xpu is not supported for torch.Generator()")
|
||||
gen = MagicMock()
|
||||
gen.manual_seed.return_value = f"gen:{device}"
|
||||
return gen
|
||||
|
||||
fake_torch = MagicMock()
|
||||
fake_torch.Generator.side_effect = fake_generator
|
||||
with patch.object(wr, "torch", fake_torch):
|
||||
assert wr._make_seed_generator("xpu", 123) == "gen:cpu"
|
||||
assert wr._make_seed_generator("cuda", 123) == "gen:cuda"
|
||||
|
||||
|
||||
class TestMpsErrorDetection:
|
||||
"""Tests for MPS error detection helper."""
|
||||
|
||||
Reference in New Issue
Block a user