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:
xchacha20-poly1305
2026-05-29 11:13:23 -07:00
committed by GitHub
co-authored by Claude Opus 4.8 Victor Kuznetsov
parent 1598c499fe
commit 0c7ff1874e
5 changed files with 114 additions and 11 deletions
+45 -2
View File
@@ -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."""