mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 09:05:17 -07:00
Merge pull request #320 from jamiepine/feat/intel-xpu-support
feat: Intel Arc (XPU) GPU support
This commit is contained in:
@@ -155,6 +155,20 @@ def _get_gpu_status() -> str:
|
|||||||
return "MPS (Apple Silicon)"
|
return "MPS (Apple Silicon)"
|
||||||
elif backend_type == "mlx":
|
elif backend_type == "mlx":
|
||||||
return "Metal (Apple Silicon via MLX)"
|
return "Metal (Apple Silicon via MLX)"
|
||||||
|
|
||||||
|
# Intel XPU (Arc / Data Center) via IPEX
|
||||||
|
try:
|
||||||
|
import intel_extension_for_pytorch # noqa: F401
|
||||||
|
|
||||||
|
if hasattr(torch, "xpu") and torch.xpu.is_available():
|
||||||
|
try:
|
||||||
|
xpu_name = torch.xpu.get_device_name(0)
|
||||||
|
except Exception:
|
||||||
|
xpu_name = "Intel GPU"
|
||||||
|
return f"XPU ({xpu_name})"
|
||||||
|
except ImportError:
|
||||||
|
pass
|
||||||
|
|
||||||
return "None (CPU only)"
|
return "None (CPU only)"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -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],
|
||||||
|
|||||||
@@ -18,6 +18,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,
|
||||||
patch_chatterbox_f32,
|
patch_chatterbox_f32,
|
||||||
@@ -48,7 +50,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 +119,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(
|
||||||
@@ -200,7 +199,7 @@ class ChatterboxTTSBackend:
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
torch.manual_seed(seed)
|
manual_seed(seed, self._device)
|
||||||
|
|
||||||
logger.info(f"[Chatterbox] Generating: lang={language}")
|
logger.info(f"[Chatterbox] Generating: lang={language}")
|
||||||
|
|
||||||
@@ -220,10 +219,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
|
||||||
|
|
||||||
|
|||||||
@@ -18,6 +18,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,
|
||||||
patch_chatterbox_f32,
|
patch_chatterbox_f32,
|
||||||
@@ -48,7 +50,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 +118,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(
|
||||||
@@ -181,7 +180,7 @@ class ChatterboxTurboTTSBackend:
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
if seed is not None:
|
if seed is not None:
|
||||||
torch.manual_seed(seed)
|
manual_seed(seed, self._device)
|
||||||
|
|
||||||
logger.info("[Chatterbox Turbo] Generating (English)")
|
logger.info("[Chatterbox Turbo] Generating (English)")
|
||||||
|
|
||||||
@@ -200,10 +199,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
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
|
|
||||||
|
|||||||
@@ -12,7 +12,14 @@ 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,
|
||||||
|
manual_seed,
|
||||||
|
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 +37,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 +76,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 +91,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")
|
||||||
|
|
||||||
@@ -154,12 +164,8 @@ class LuxTTSBackend:
|
|||||||
await self.load_model()
|
await self.load_model()
|
||||||
|
|
||||||
def _generate_sync():
|
def _generate_sync():
|
||||||
import torch
|
|
||||||
|
|
||||||
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)
|
|
||||||
|
|
||||||
wav = self.model.generate_speech(
|
wav = self.model.generate_speech(
|
||||||
text=text,
|
text=text,
|
||||||
|
|||||||
@@ -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,
|
||||||
)
|
)
|
||||||
@@ -122,8 +124,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")
|
||||||
|
|
||||||
@@ -215,9 +216,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(
|
||||||
@@ -300,8 +299,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")
|
||||||
|
|
||||||
|
|||||||
@@ -110,6 +110,11 @@ async def health():
|
|||||||
vram_used = None
|
vram_used = None
|
||||||
if has_cuda:
|
if has_cuda:
|
||||||
vram_used = torch.cuda.memory_allocated() / 1024 / 1024
|
vram_used = torch.cuda.memory_allocated() / 1024 / 1024
|
||||||
|
elif has_xpu:
|
||||||
|
try:
|
||||||
|
vram_used = torch.xpu.memory_allocated() / 1024 / 1024
|
||||||
|
except Exception:
|
||||||
|
pass # memory_allocated() may not be available on all IPEX versions
|
||||||
|
|
||||||
model_loaded = False
|
model_loaded = False
|
||||||
model_size = None
|
model_size = None
|
||||||
@@ -162,7 +167,10 @@ async def health():
|
|||||||
gpu_type=gpu_type,
|
gpu_type=gpu_type,
|
||||||
vram_used_mb=vram_used,
|
vram_used_mb=vram_used,
|
||||||
backend_type=backend_type,
|
backend_type=backend_type,
|
||||||
backend_variant=os.environ.get("VOICEBOX_BACKEND_VARIANT", "cuda" if torch.cuda.is_available() else "cpu"),
|
backend_variant=os.environ.get(
|
||||||
|
"VOICEBOX_BACKEND_VARIANT",
|
||||||
|
"cuda" if torch.cuda.is_available() else ("xpu" if has_xpu else "cpu"),
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -69,10 +69,22 @@ 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' })
|
$gpus = Get-CimInstance Win32_VideoController | Select-Object -ExpandProperty Name
|
||||||
|
Write-Host "Detected GPUs: $($gpus -join ', ')"
|
||||||
|
$hasNvidia = ($gpus | Where-Object { $_ -match 'NVIDIA' }).Count -gt 0
|
||||||
|
$hasIntelArc = ($gpus | Where-Object { $_ -match 'Arc' }).Count -gt 0
|
||||||
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; \
|
||||||
|
} else { \
|
||||||
|
Write-Host "No NVIDIA or Intel Arc GPU detected — using CPU-only PyTorch."; \
|
||||||
|
Write-Host "If you have an Intel Arc GPU, install XPU support manually:"; \
|
||||||
|
Write-Host " pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/xpu"; \
|
||||||
|
Write-Host " 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
|
||||||
|
|||||||
Reference in New Issue
Block a user