studio: extend the _grouped_mm null-kernel guard to Linux ROCm RDNA4 (gfx1201) (#7292)

* studio: extend the _grouped_mm null-kernel guard to Linux ROCm RDNA4

torch._grouped_mm has a null HIP kernel on RDNA4 (gfx1200/gfx1201) at
ROCm <= 7.12 (fixed in 7.13; ROCm/TheRock #5284). The existing guard that
registers a Python mm/bmm fallback was win32-only, so Linux gfx1201 (e.g.
R9700 Pro on Ubuntu) hits the null kernel -> illegal instruction during
training.

Extract the fallback registration into a module-level helper
(_install_grouped_mm_cpu_fallback) and add a Linux branch that installs it,
gated on gfx1200/gfx1201 AND HIP < 7.13 so NVIDIA/CUDA and every non-RDNA4
AMD arch are untouched, and it is a no-op on fixed runtimes. The Windows
path now calls the same helper with identical behavior.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Resolve HIP version from torch.__version__ when version.hip is unset for PR #7292

AMD SDK / Radeon ROCm wheels leave torch.version.hip empty and encode the
version only in torch.__version__ (e.g. +rocm7.12). The Linux gfx120X guard
parsed version.hip only, so those affected installs skipped the fallback and
still hit the null _grouped_mm kernel. Mirror the Windows parse: version.hip,
then the embedded rocmX.Y, then assume affected unless a post-fix rocmsdk wheel.

* Scan all GPUs and add RDNA4 name fallback for _grouped_mm guard in PR #7292

* worker.py: tighten gfx120X Linux guard comments (no code change)

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
Daniel Han 2026-07-21 18:01:07 -07:00 committed by GitHub
commit 207a9f00bf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194

View file

@ -90,6 +90,79 @@ _FAST_PATH_HOOKS_SKIP_ENV = "UNSLOTH_STUDIO_SKIP_FAST_PATH_HOOKS"
# run_training_process() and isn't GC'd mid-run.
_WINDOWS_ROCM_GROUPED_MM_LIB = None
def _install_grouped_mm_cpu_fallback(torch_mod, logger, label):
"""Register a Python mm/bmm fallback for torch._grouped_mm and return the Library.
RDNA4 (gfx1200/gfx1201) ships a null HIP _grouped_mm kernel on ROCm <= 7.12
(fixed in 7.13; ROCm/TheRock #5284). JitDecomp dispatches _grouped_mm to the
null kernel and crashes; overriding the CUDA dispatch key bypasses it. Shared
by the Windows and Linux ROCm guards. Keep the returned Library referenced so
the registration outlives the caller.
"""
import warnings as _warnings
_gm_lib = torch_mod.library.Library("aten", "IMPL")
def _grouped_mm_safe_impl(
self,
mat2,
offs = None,
bias = None,
out_dtype = None,
):
"""Python mm/bmm fallback for _grouped_mm on gfx120X (null HIP kernel, ROCm <= 7.12)."""
_t = torch_mod
if offs is None:
# No offsets: 2-D -> mm, 3-D batched -> bmm (unconditional mm broke 3-D MoE).
if self.dim() == 3 and mat2.dim() == 3:
result = _t.bmm(self.contiguous(), mat2.contiguous())
elif self.dim() == 3 and mat2.dim() == 2:
result = _t.matmul(self.contiguous(), mat2.contiguous())
elif self.dim() == 2 and mat2.dim() == 3:
result = _t.matmul(self.contiguous(), mat2.contiguous())
else:
result = _t.mm(self.contiguous(), mat2.contiguous())
else:
# Grouped: offs[i] is the exclusive end-row of group i.
offs_list = offs.tolist()
pieces = []
prev = 0
for idx, end in enumerate(offs_list):
end = int(end)
a_part = self[prev:end].contiguous()
b_part = mat2[idx].contiguous() if mat2.dim() == 3 else mat2.contiguous()
pieces.append(_t.mm(a_part, b_part))
prev = end
# Include trailing rows not covered by offs.
if prev < self.shape[0]:
a_tail = self[prev:].contiguous()
b_tail = mat2[-1].contiguous() if mat2.dim() == 3 else mat2.contiguous()
pieces.append(_t.mm(a_tail, b_tail))
result = (
_t.cat(pieces, dim = 0)
if pieces
else _t.zeros(0, mat2.shape[-1], device = self.device, dtype = self.dtype)
)
if bias is not None:
result = result + bias
if out_dtype is not None:
result = result.to(out_dtype)
elif result.dtype != self.dtype:
result = result.to(self.dtype)
return result
with _warnings.catch_warnings():
_warnings.simplefilter("ignore")
_gm_lib.impl("_grouped_mm", _grouped_mm_safe_impl, "CUDA")
logger.info(
"%s: patched _grouped_mm CUDA dispatch (null HIP kernel on gfx120X, "
"ROCm <= 7.12 -- bypassed with Python mm fallback)",
label,
)
return _gm_lib
# Subprocesses don't inherit os.add_dll_directory registrations. Replicate
# main.py's Windows ROCm DLL setup so the first `import torch` finds
# amdhip64.dll. Handles retained at module scope so they aren't GC'd.
@ -2689,80 +2762,8 @@ def run_training_process(*, event_queue: Any, stop_queue: Any, config: dict) ->
# so 7.13+ uses the real GPU kernel.
if not _hip_ver_at_least(7, 13):
try:
import warnings as _warnings
_gm_lib = _torch_for_rocm.library.Library("aten", "IMPL")
def _grouped_mm_safe_impl(
self,
mat2,
offs = None,
bias = None,
out_dtype = None,
):
"""Python mm/bmm fallback for _grouped_mm on gfx1200 (null HIP kernel, ROCm ≤ 7.12)."""
_t = _torch_for_rocm
if offs is None:
# No offsets: 2-D -> mm, 3-D batched -> bmm
# (unconditional mm broke 3-D MoE).
if self.dim() == 3 and mat2.dim() == 3:
result = _t.bmm(self.contiguous(), mat2.contiguous())
elif self.dim() == 3 and mat2.dim() == 2:
# Broadcast 2-D mat2 across the batch dim.
result = _t.matmul(self.contiguous(), mat2.contiguous())
elif self.dim() == 2 and mat2.dim() == 3:
# Broadcast 2-D self across batch via matmul.
result = _t.matmul(self.contiguous(), mat2.contiguous())
else:
result = _t.mm(self.contiguous(), mat2.contiguous())
else:
# Grouped: offs[i] is the exclusive end-row of group i.
offs_list = offs.tolist()
pieces = []
prev = 0
for idx, end in enumerate(offs_list):
end = int(end)
a_part = self[prev:end].contiguous()
if mat2.dim() == 3:
b_part = mat2[idx].contiguous()
else:
b_part = mat2.contiguous()
pieces.append(_t.mm(a_part, b_part))
prev = end
# Include trailing rows not covered by offs.
if prev < self.shape[0]:
a_tail = self[prev:].contiguous()
b_tail = (
mat2[-1].contiguous() if mat2.dim() == 3 else mat2.contiguous()
)
pieces.append(_t.mm(a_tail, b_tail))
result = (
_t.cat(pieces, dim = 0)
if pieces
else _t.zeros(
0,
mat2.shape[-1],
device = self.device,
dtype = self.dtype,
)
)
if bias is not None:
result = result + bias
if out_dtype is not None:
result = result.to(out_dtype)
elif result.dtype != self.dtype:
result = result.to(self.dtype)
return result
with _warnings.catch_warnings():
_warnings.simplefilter("ignore")
_gm_lib.impl("_grouped_mm", _grouped_mm_safe_impl, "CUDA")
_WINDOWS_ROCM_GROUPED_MM_LIB = _gm_lib # prevent GC
logger.info(
"Windows ROCm: patched _grouped_mm CUDA dispatch "
"(null HIP kernel on gfx1200, ROCm ≤ 7.12 — "
"bypassed with Python mm fallback)"
_WINDOWS_ROCM_GROUPED_MM_LIB = _install_grouped_mm_cpu_fallback(
_torch_for_rocm, logger, "Windows ROCm"
)
except Exception as _patch_exc:
logger.warning(
@ -2776,6 +2777,44 @@ def run_training_process(*, event_queue: Any, stop_queue: Any, config: dict) ->
"skipping Python fallback (AMD fixed gfx1200 null kernel in ROCm 7.13)"
)
# ── 1f-linux. Linux ROCm RDNA4 _grouped_mm null kernel ──
# The win32 guard above misses Linux: RDNA4 (gfx1200/gfx1201) hits the same null
# HIP _grouped_mm kernel at ROCm <= 7.12 (fixed 7.13, ROCm/TheRock #5284). Gate on
# arch + HIP < 7.13 so NVIDIA/CUDA and non-RDNA4 AMD are untouched; no-op if fixed.
if sys.platform.startswith("linux") and _hw.IS_ROCM:
try:
_torch_lin = sys.modules.get("torch")
if _torch_lin is not None and _torch_lin.cuda.is_available():
# Prefer torch.version.hip, else rocmX.Y from torch.__version__ (AMD
# SDK / Radeon wheels leave version.hip unset). Unknown version on a
# gfx120X build -> assume affected unless it is a post-fix rocmsdk wheel.
_hip_str = str(getattr(getattr(_torch_lin, "version", None), "hip", "") or "")
_ver = getattr(_torch_lin, "__version__", "").lower()
_m = re.match(r"(\d+)\.(\d+)", _hip_str) or re.search(r"rocm(\d+)\.(\d+)", _ver)
if _m:
_hip_lt_713 = (int(_m.group(1)), int(_m.group(2))) < (7, 13)
else:
_hip_lt_713 = "rocmsdk" not in _ver
# Scan every visible GPU (device_map="balanced" can place layers on a
# later RDNA4 card, so device 0 is not enough). Match gfx120X by arch,
# or by RX 9000 / R9700 name when the wheel omits gcnArchName.
_rdna4 = False
for _i in range(_torch_lin.cuda.device_count()):
_props = _torch_lin.cuda.get_device_properties(_i)
_lin_arch, _ = _rocm_classify_unified_memory(_props)
_lin_name = (getattr(_props, "name", "") or "").lower()
if _lin_arch.lower() in ("gfx1200", "gfx1201") or (
not _lin_arch and re.search(r"rx\s*90[0-9]0|r9700", _lin_name)
):
_rdna4 = True
break
if _rdna4 and _hip_lt_713:
_WINDOWS_ROCM_GROUPED_MM_LIB = _install_grouped_mm_cpu_fallback(
_torch_lin, logger, "Linux ROCm gfx120X"
)
except Exception as _gm_lin_exc:
logger.warning("Linux ROCm gfx120X: could not patch _grouped_mm: %s", _gm_lin_exc)
# ── 1g. ROCm OOM guard ──
# On ROCm, exhausting VRAM can hang the HIP driver instead of raising.
# set_per_process_memory_fraction caps the allocator so PyTorch raises