feat: add Intel Arc (XPU) GPU support across all backends

Auto-detect Intel Arc GPUs during Windows setup and install PyTorch
with XPU support + intel-extension-for-pytorch. Enable allow_xpu=True
on all TTS backends (Chatterbox, Chatterbox Turbo, Hume TADA, LuxTTS)
that previously only supported CUDA. Add shared empty_device_cache()
and manual_seed() helpers in base.py to handle XPU memory management
and reproducible seeding alongside CUDA.
This commit is contained in:
James Pine
2026-03-18 11:24:51 -07:00
parent ffc1b54812
commit 83ebababe7
7 changed files with 83 additions and 51 deletions
+31
View File
@@ -126,6 +126,37 @@ def get_torch_device(
return "cpu" return "cpu"
def empty_device_cache(device: str) -> None:
"""
Free cached memory on the given device (CUDA or XPU).
Backends should call this after unloading models so VRAM is returned
to the OS.
"""
import torch
if device == "cuda" and torch.cuda.is_available():
torch.cuda.empty_cache()
elif device == "xpu" and hasattr(torch, "xpu"):
torch.xpu.empty_cache()
def manual_seed(seed: int, device: str) -> None:
"""
Set the random seed on both CPU and the active accelerator.
Covers CUDA and Intel XPU so that generation is reproducible
regardless of which GPU backend is in use.
"""
import torch
torch.manual_seed(seed)
if device == "cuda" and torch.cuda.is_available():
torch.cuda.manual_seed(seed)
elif device == "xpu" and hasattr(torch, "xpu"):
torch.xpu.manual_seed(seed)
async def combine_voice_prompts( async def combine_voice_prompts(
audio_paths: List[str], audio_paths: List[str],
reference_texts: List[str], reference_texts: List[str],
+4 -9
View File
@@ -18,6 +18,7 @@ from . import TTSBackend
from .base import ( from .base import (
is_model_cached, is_model_cached,
get_torch_device, get_torch_device,
empty_device_cache,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
model_load_progress, model_load_progress,
patch_chatterbox_f32, patch_chatterbox_f32,
@@ -48,7 +49,7 @@ class ChatterboxTTSBackend:
self._model_load_lock = asyncio.Lock() self._model_load_lock = asyncio.Lock()
def _get_device(self) -> str: def _get_device(self) -> str:
return get_torch_device(force_cpu_on_mac=True) return get_torch_device(force_cpu_on_mac=True, allow_xpu=True)
def is_loaded(self) -> bool: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -117,10 +118,7 @@ class ChatterboxTTSBackend:
del self.model del self.model
self.model = None self.model = None
self._device = None self._device = None
if device == "cuda": empty_device_cache(device)
import torch
torch.cuda.empty_cache()
logger.info("Chatterbox unloaded") logger.info("Chatterbox unloaded")
async def create_voice_prompt( async def create_voice_prompt(
@@ -220,10 +218,7 @@ class ChatterboxTTSBackend:
else: else:
audio = np.asarray(wav, dtype=np.float32) audio = np.asarray(wav, dtype=np.float32)
sample_rate = ( sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000)
getattr(self.model, "sr", None)
or getattr(self.model, "sample_rate", 24000)
)
return audio, sample_rate return audio, sample_rate
+4 -9
View File
@@ -18,6 +18,7 @@ from . import TTSBackend
from .base import ( from .base import (
is_model_cached, is_model_cached,
get_torch_device, get_torch_device,
empty_device_cache,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
model_load_progress, model_load_progress,
patch_chatterbox_f32, patch_chatterbox_f32,
@@ -48,7 +49,7 @@ class ChatterboxTurboTTSBackend:
self._model_load_lock = asyncio.Lock() self._model_load_lock = asyncio.Lock()
def _get_device(self) -> str: def _get_device(self) -> str:
return get_torch_device(force_cpu_on_mac=True) return get_torch_device(force_cpu_on_mac=True, allow_xpu=True)
def is_loaded(self) -> bool: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -116,10 +117,7 @@ class ChatterboxTurboTTSBackend:
del self.model del self.model
self.model = None self.model = None
self._device = None self._device = None
if device == "cuda": empty_device_cache(device)
import torch
torch.cuda.empty_cache()
logger.info("Chatterbox Turbo unloaded") logger.info("Chatterbox Turbo unloaded")
async def create_voice_prompt( async def create_voice_prompt(
@@ -200,10 +198,7 @@ class ChatterboxTurboTTSBackend:
else: else:
audio = np.asarray(wav, dtype=np.float32) audio = np.asarray(wav, dtype=np.float32)
sample_rate = ( sample_rate = getattr(self.model, "sr", None) or getattr(self.model, "sample_rate", 24000)
getattr(self.model, "sr", None)
or getattr(self.model, "sample_rate", 24000)
)
return audio, sample_rate return audio, sample_rate
+19 -20
View File
@@ -24,6 +24,8 @@ from . import TTSBackend
from .base import ( from .base import (
is_model_cached, is_model_cached,
get_torch_device, get_torch_device,
empty_device_cache,
manual_seed,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
model_load_progress, model_load_progress,
) )
@@ -66,7 +68,7 @@ class HumeTadaBackend:
def _get_device(self) -> str: def _get_device(self) -> str:
# Force CPU on macOS — MPS has issues with flow matching # Force CPU on macOS — MPS has issues with flow matching
# and large vocab lm_head (>65536 output channels) # and large vocab lm_head (>65536 output channels)
return get_torch_device(force_cpu_on_mac=True) return get_torch_device(force_cpu_on_mac=True, allow_xpu=True)
def is_loaded(self) -> bool: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -105,6 +107,7 @@ class HumeTadaBackend:
# package. The real package pulls in onnx/tensorboard/matplotlib via # package. The real package pulls in onnx/tensorboard/matplotlib via
# descript-audiotools, so we use a lightweight shim instead. # descript-audiotools, so we use a lightweight shim instead.
from ..utils.dac_shim import install_dac_shim from ..utils.dac_shim import install_dac_shim
install_dac_shim() install_dac_shim()
import torch import torch
@@ -142,9 +145,12 @@ class HumeTadaBackend:
allow_patterns=["tokenizer*", "special_tokens*"], allow_patterns=["tokenizer*", "special_tokens*"],
) )
# Determine dtype — use bf16 on CUDA for ~50% memory savings # Determine dtype — use bf16 on CUDA/XPU for ~50% memory savings
if device == "cuda" and torch.cuda.is_bf16_supported(): if device == "cuda" and torch.cuda.is_bf16_supported():
model_dtype = torch.bfloat16 model_dtype = torch.bfloat16
elif device == "xpu":
# Intel Arc (Alchemist+) supports bf16 natively
model_dtype = torch.bfloat16
else: else:
model_dtype = torch.float32 model_dtype = torch.float32
@@ -153,14 +159,14 @@ class HumeTadaBackend:
# This avoids monkey-patching AutoTokenizer.from_pretrained # This avoids monkey-patching AutoTokenizer.from_pretrained
# which corrupts the classmethod descriptor for other engines. # which corrupts the classmethod descriptor for other engines.
from tada.modules.aligner import AlignerConfig from tada.modules.aligner import AlignerConfig
AlignerConfig.tokenizer_name = tokenizer_path AlignerConfig.tokenizer_name = tokenizer_path
# Load encoder (only needed for voice prompt encoding) # Load encoder (only needed for voice prompt encoding)
from tada.modules.encoder import Encoder from tada.modules.encoder import Encoder
logger.info("Loading TADA encoder...") logger.info("Loading TADA encoder...")
self.encoder = Encoder.from_pretrained( self.encoder = Encoder.from_pretrained(TADA_CODEC_REPO, subfolder="encoder").to(device)
TADA_CODEC_REPO, subfolder="encoder"
).to(device)
self.encoder.eval() self.encoder.eval()
# Load the causal LM (includes decoder for wav generation). # Load the causal LM (includes decoder for wav generation).
@@ -169,12 +175,11 @@ class HumeTadaBackend:
# which hits the gated repo. Pre-load the config from HF, # which hits the gated repo. Pre-load the config from HF,
# inject the local tokenizer path, then pass it in. # inject the local tokenizer path, then pass it in.
from tada.modules.tada import TadaForCausalLM, TadaConfig from tada.modules.tada import TadaForCausalLM, TadaConfig
logger.info(f"Loading TADA {model_size} model...") logger.info(f"Loading TADA {model_size} model...")
config = TadaConfig.from_pretrained(repo) config = TadaConfig.from_pretrained(repo)
config.tokenizer_name = tokenizer_path config.tokenizer_name = tokenizer_path
self.model = TadaForCausalLM.from_pretrained( self.model = TadaForCausalLM.from_pretrained(repo, config=config, torch_dtype=model_dtype).to(device)
repo, config=config, torch_dtype=model_dtype
).to(device)
self.model.eval() self.model.eval()
logger.info(f"HumeAI TADA {model_size} loaded successfully on {device}") logger.info(f"HumeAI TADA {model_size} loaded successfully on {device}")
@@ -188,11 +193,11 @@ class HumeTadaBackend:
del self.encoder del self.encoder
self.encoder = None self.encoder = None
device = self._device
self._device = None self._device = None
import torch if device:
if torch.cuda.is_available(): empty_device_cache(device)
torch.cuda.empty_cache()
logger.info("HumeAI TADA unloaded") logger.info("HumeAI TADA unloaded")
@@ -213,9 +218,7 @@ class HumeTadaBackend:
""" """
await self.load_model(self.model_size) await self.load_model(self.model_size)
cache_key = ( cache_key = ("tada_" + get_cache_key(audio_path, reference_text)) if use_cache else None
"tada_" + get_cache_key(audio_path, reference_text)
) if use_cache else None
if cache_key: if cache_key:
cached = get_cached_voice_prompt(cache_key) cached = get_cached_voice_prompt(cache_key)
@@ -239,9 +242,7 @@ class HumeTadaBackend:
# Encode with forced alignment # Encode with forced alignment
text_arg = [reference_text] if reference_text else None text_arg = [reference_text] if reference_text else None
prompt = self.encoder( prompt = self.encoder(audio, text=text_arg, sample_rate=sr)
audio, text=text_arg, sample_rate=sr
)
# Serialize EncoderOutput to a dict of CPU tensors for caching # Serialize EncoderOutput to a dict of CPU tensors for caching
prompt_dict = {} prompt_dict = {}
@@ -299,9 +300,7 @@ class HumeTadaBackend:
from tada.modules.encoder import EncoderOutput from tada.modules.encoder import EncoderOutput
if seed is not None: if seed is not None:
torch.manual_seed(seed) manual_seed(seed, self._device)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
device = self._device device = self._device
+15 -6
View File
@@ -12,7 +12,13 @@ from typing import Optional, Tuple
import numpy as np import numpy as np
from . import TTSBackend from . import TTSBackend
from .base import is_model_cached, get_torch_device, combine_voice_prompts as _combine_voice_prompts, model_load_progress from .base import (
is_model_cached,
get_torch_device,
empty_device_cache,
combine_voice_prompts as _combine_voice_prompts,
model_load_progress,
)
from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt from ..utils.cache import get_cache_key, get_cached_voice_prompt, cache_voice_prompt
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
@@ -30,7 +36,7 @@ class LuxTTSBackend:
self._device = None self._device = None
def _get_device(self) -> str: def _get_device(self) -> str:
return get_torch_device(allow_mps=True) return get_torch_device(allow_mps=True, allow_xpu=True)
def is_loaded(self) -> bool: def is_loaded(self) -> bool:
return self.model is not None return self.model is not None
@@ -69,9 +75,12 @@ class LuxTTSBackend:
if device == "cpu": if device == "cpu":
import os import os
threads = os.cpu_count() or 4 threads = os.cpu_count() or 4
self.model = LuxTTS( self.model = LuxTTS(
model_path=LUXTTS_HF_REPO, device="cpu", threads=min(threads, 8), model_path=LUXTTS_HF_REPO,
device="cpu",
threads=min(threads, 8),
) )
else: else:
self.model = LuxTTS(model_path=LUXTTS_HF_REPO, device=device) self.model = LuxTTS(model_path=LUXTTS_HF_REPO, device=device)
@@ -81,12 +90,12 @@ class LuxTTSBackend:
def unload_model(self) -> None: def unload_model(self) -> None:
"""Unload model to free memory.""" """Unload model to free memory."""
if self.model is not None: if self.model is not None:
device = self.device
del self.model del self.model
self.model = None self.model = None
self._device = None
import torch empty_device_cache(device)
if torch.cuda.is_available():
torch.cuda.empty_cache()
logger.info("LuxTTS unloaded") logger.info("LuxTTS unloaded")
+5 -7
View File
@@ -14,6 +14,8 @@ from . import TTSBackend, STTBackend, LANGUAGE_CODE_TO_NAME, WHISPER_HF_REPOS
from .base import ( from .base import (
is_model_cached, is_model_cached,
get_torch_device, get_torch_device,
empty_device_cache,
manual_seed,
combine_voice_prompts as _combine_voice_prompts, combine_voice_prompts as _combine_voice_prompts,
model_load_progress, model_load_progress,
) )
@@ -120,8 +122,7 @@ class PyTorchTTSBackend:
self.model = None self.model = None
self._current_model_size = None self._current_model_size = None
if torch.cuda.is_available(): empty_device_cache(self.device)
torch.cuda.empty_cache()
logger.info("TTS model unloaded") logger.info("TTS model unloaded")
@@ -213,9 +214,7 @@ class PyTorchTTSBackend:
"""Run synchronous generation in thread pool.""" """Run synchronous generation in thread pool."""
# Set seed if provided # Set seed if provided
if seed is not None: if seed is not None:
torch.manual_seed(seed) manual_seed(seed, self.device)
if torch.cuda.is_available():
torch.cuda.manual_seed(seed)
# Generate audio - this is the blocking operation # Generate audio - this is the blocking operation
wavs, sample_rate = self.model.generate_voice_clone( wavs, sample_rate = self.model.generate_voice_clone(
@@ -297,8 +296,7 @@ class PyTorchSTTBackend:
self.model = None self.model = None
self.processor = None self.processor = None
if torch.cuda.is_available(): empty_device_cache(self.device)
torch.cuda.empty_cache()
logger.info("Whisper model unloaded") logger.info("Whisper model unloaded")
+5
View File
@@ -70,9 +70,14 @@ setup-python:
Write-Host "Installing Python dependencies..." Write-Host "Installing Python dependencies..."
& "{{ python }}" -m pip install --upgrade pip -q & "{{ python }}" -m pip install --upgrade pip -q
$hasNvidia = $null -ne (Get-WmiObject Win32_VideoController | Where-Object { $_.Name -match 'NVIDIA' }) $hasNvidia = $null -ne (Get-WmiObject Win32_VideoController | Where-Object { $_.Name -match 'NVIDIA' })
$hasIntelArc = $null -ne (Get-WmiObject Win32_VideoController | Where-Object { $_.Name -match 'Intel.*Arc' })
if ($hasNvidia) { \ if ($hasNvidia) { \
Write-Host "NVIDIA GPU detected — installing PyTorch with CUDA support..."; \ Write-Host "NVIDIA GPU detected — installing PyTorch with CUDA support..."; \
& "{{ pip }}" install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128; \ & "{{ pip }}" install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128; \
} elseif ($hasIntelArc) { \
Write-Host "Intel Arc GPU detected — installing PyTorch with XPU support..."; \
& "{{ pip }}" install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/xpu; \
& "{{ pip }}" install intel-extension-for-pytorch --index-url https://download.pytorch.org/whl/xpu; \
} }
& "{{ pip }}" install -r {{ backend_dir }}/requirements.txt & "{{ pip }}" install -r {{ backend_dir }}/requirements.txt
& "{{ pip }}" install --no-deps chatterbox-tts & "{{ pip }}" install --no-deps chatterbox-tts