mirror of
https://github.com/jamiepine/voicebox.git
synced 2026-10-03 00:55:14 -07:00
fix(backend): detect AMD GPU before setting HSA_OVERRIDE_GFX_VERSION (#785)
Previously, HSA_OVERRIDE_GFX_VERSION=10.3.0 was unconditionally set for all AMD GPUs, which caused suboptimal performance on RDNA 3/4 GPUs (gfx11xx/gfx12xx) that have native ROCm support. Now uses rocminfo to detect all GPUs and only sets the override for systems where the oldest GPU needs it (RDNA 2 and older, gfx10xx and below). Newer GPUs are left untouched. Addresses CodeRabbit review: - Case-insensitive regex matching on lowercased line - Log level changed to INFO for rocminfo failures - Multi-GPU support: iterates all GPUs, uses oldest for decision Fixes #469 Signed-off-by: Amitesh Gupta Signed-off-by: Amitesh Gupta Signed-off-by: singlaamitesh <[email protected]>
This commit is contained in:
+54
-1
@@ -3,6 +3,8 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
|
import subprocess
|
||||||
import sys
|
import sys
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
@@ -37,8 +39,59 @@ logging.basicConfig(
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# AMD GPU environment variables must be set before torch import
|
# AMD GPU environment variables must be set before torch import
|
||||||
|
# Only set HSA_OVERRIDE_GFX_VERSION for older GPUs that need it.
|
||||||
|
# RDNA 3+ (gfx1100+) and RDNA 4 (gfx1200+) are natively supported by ROCm
|
||||||
|
# and the override can cause suboptimal performance or errors.
|
||||||
if not os.environ.get("HSA_OVERRIDE_GFX_VERSION"):
|
if not os.environ.get("HSA_OVERRIDE_GFX_VERSION"):
|
||||||
os.environ["HSA_OVERRIDE_GFX_VERSION"] = "10.3.0"
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["rocminfo"],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=5,
|
||||||
|
)
|
||||||
|
if result.returncode == 0:
|
||||||
|
# Collect all GPUs found in rocminfo output
|
||||||
|
gfx_versions = []
|
||||||
|
for line in result.stdout.splitlines():
|
||||||
|
line_lower = line.lower()
|
||||||
|
if "gfx" in line_lower:
|
||||||
|
match = re.search(r"(gfx\d+)", line_lower)
|
||||||
|
if match:
|
||||||
|
gfx_versions.append(match.group(1))
|
||||||
|
|
||||||
|
if gfx_versions:
|
||||||
|
# Check if any GPU needs the override (RDNA 2 and older)
|
||||||
|
# Use the oldest GPU (lowest gfx number) for the decision
|
||||||
|
try:
|
||||||
|
gfx_nums = []
|
||||||
|
for v in gfx_versions:
|
||||||
|
m = re.search(r"\d+", v)
|
||||||
|
if m:
|
||||||
|
gfx_nums.append(int(m.group()))
|
||||||
|
if gfx_nums:
|
||||||
|
oldest_num = min(gfx_nums)
|
||||||
|
oldest_gfx = gfx_versions[gfx_nums.index(oldest_num)]
|
||||||
|
if oldest_num < 1100:
|
||||||
|
os.environ["HSA_OVERRIDE_GFX_VERSION"] = "10.3.0"
|
||||||
|
logger.info(
|
||||||
|
"AMD GPU detected (%s), setting HSA_OVERRIDE_GFX_VERSION=10.3.0 for compatibility. All GPUs: %s",
|
||||||
|
oldest_gfx,
|
||||||
|
", ".join(gfx_versions),
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.info(
|
||||||
|
"AMD GPU detected (%s), native ROCm support available, skipping HSA_OVERRIDE_GFX_VERSION. All GPUs: %s",
|
||||||
|
oldest_gfx,
|
||||||
|
", ".join(gfx_versions),
|
||||||
|
)
|
||||||
|
except (ValueError, AttributeError) as e:
|
||||||
|
logger.info("Could not parse GPU version from rocminfo output: %s", e)
|
||||||
|
except (FileNotFoundError, subprocess.TimeoutExpired, Exception) as e:
|
||||||
|
logger.info(
|
||||||
|
"Could not detect AMD GPU via rocminfo, skipping automatic HSA_OVERRIDE_GFX_VERSION configuration: %s",
|
||||||
|
e,
|
||||||
|
)
|
||||||
if not os.environ.get("MIOPEN_LOG_LEVEL"):
|
if not os.environ.get("MIOPEN_LOG_LEVEL"):
|
||||||
os.environ["MIOPEN_LOG_LEVEL"] = "4"
|
os.environ["MIOPEN_LOG_LEVEL"] = "4"
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user