Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions projects/hipblaslt/cmake/hipblaslt_python.cmake
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ function(hipblaslt_configure_bundled_python_command python_binary_dir asan_optio
"${CMAKE_COMMAND}" -E env
"PYTHONPATH=${_python_path}"
"PATH=${_path}"
"ROCM_VERSION=${hip_VERSION}"
"${asan_options}"
--
"${Python3_EXECUTABLE}"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -81,10 +81,7 @@
from Tensile.verify_stinky_comment_vs_elf_text import verify_stinky_paths
from Tensile.Toolchain.Assembly import makeAssemblyToolchain, buildAssemblyCodeObjectFiles
from Tensile.Toolchain.Source import makeSourceToolchain, buildSourceCodeObjectFiles
from Tensile.Toolchain.Validators import (
ToolchainDefaults,
validateToolchain,
)
from Tensile.Toolchain.Validators import validateToolchain
from Tensile.Toolchain.Component import Assembler
from Tensile.Utilities.Decorators.Profile import profile
from Tensile.Utilities.Decorators.Timing import timing
Expand Down Expand Up @@ -1086,12 +1083,11 @@ def run():
arguments = parseArguments()
setVerbosity(arguments["PrintLevel"])
outputPath = Path(ensurePath(os.path.abspath(arguments["OutputPath"])))
cxxCompiler, _, offloadBundler, _, _ = validateToolchain(
cxxCompiler, _, offloadBundler, _ = validateToolchain(
arguments["CxxCompiler"],
arguments["CCompiler"],
arguments["OffloadBundler"],
arguments["Assembler"],
ToolchainDefaults.HIP_CONFIG,
)

if ";" in arguments["Architecture"]:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -650,11 +650,12 @@ def _run_with_stubs(self, logic_dir, output_dir, extra_args=None):
base_args.update(extra_args)

_, info_map = _make_isa_info_map("gfx942")
validate_toolchain = MagicMock(return_value=(
"/fake/hipcc", None, "/fake/bundler", None))

with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
patch.object(M, "validateToolchain", validate_toolchain), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand All @@ -672,6 +673,13 @@ def _run_with_stubs(self, logic_dir, output_dir, extra_args=None):
patch.object(M.shutil, "rmtree"):
M.run()

validate_toolchain.assert_called_once_with(
base_args["CxxCompiler"],
base_args["CCompiler"],
base_args["OffloadBundler"],
base_args["Assembler"],
)

def test_run_completes_without_exception(self, logic_dir, output_dir):
"""Lines 881-1086: run() must complete without raising."""
self._run_with_stubs(logic_dir, output_dir)
Expand All @@ -694,7 +702,7 @@ def _spy_write(name, data, *a, **kw):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand Down Expand Up @@ -733,7 +741,7 @@ def _spy_write(name, data, *a, **kw):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand Down Expand Up @@ -773,7 +781,7 @@ def _spy_write(name, data, fmt=None, *a, **kw):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand Down Expand Up @@ -802,7 +810,7 @@ def test_run_logic_path_not_exist_exits(self, tmp_path):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand All @@ -819,7 +827,7 @@ def test_run_semicolon_arch_is_split(self, logic_dir, output_dir):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand Down Expand Up @@ -851,7 +859,7 @@ def test_run_keep_build_tmp_suppresses_rmtree(self, logic_dir, output_dir):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand Down Expand Up @@ -909,7 +917,7 @@ def _spy_write(name, data, fmt=None, *a, **kw):
with patch.object(M, "parseArguments", return_value=base_args), \
patch.object(M, "setVerbosity"), \
patch.object(M, "validateToolchain", return_value=(
"/fake/hipcc", None, "/fake/bundler", None, None)), \
"/fake/hipcc", None, "/fake/bundler", None)), \
patch.object(M, "makeIsaInfoMap", return_value=info_map), \
patch.object(M, "assignGlobalParameters"), \
patch.object(M, "makeAssemblyToolchain", return_value=MagicMock()), \
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -79,13 +79,38 @@ class _R:
C._getVersion("amdclang++", "--version", r"version\s+([\d.]+)")


def test_get_rocm_version_uses_hipconfig(monkeypatch):
@pytest.mark.parametrize(
"version, expected",
[
("6.4.43482", SemanticVersion(6, 4, 43482)),
("7.1.25424-4179531dcd", SemanticVersion(7, 1, 25424)),
("10.1.0a20260813", SemanticVersion(10, 1, 0)),
],
)
def test_get_rocm_version_uses_environment(monkeypatch, version, expected):
monkeypatch.setenv("ROCM_VERSION", version)
monkeypatch.setattr(
C,
"_getVersion",
lambda *args, **kwargs: pytest.fail("hipconfig fallback must not run"),
)
assert C.get_rocm_version() == expected


def test_get_rocm_version_rejects_invalid_environment(monkeypatch):
monkeypatch.setenv("ROCM_VERSION", "not-a-version")
with pytest.raises(RuntimeError, match="Invalid ROCM_VERSION"):
C.get_rocm_version()


def test_get_rocm_version_falls_back_to_hipconfig(monkeypatch):
seen = {}

def _fake(exe, flag, regex):
seen["exe"], seen["flag"] = exe, flag
return SemanticVersion(6, 4, 0)

monkeypatch.delenv("ROCM_VERSION", raising=False)
monkeypatch.setattr(C, "_getVersion", _fake)
assert C.get_rocm_version() == SemanticVersion(6, 4, 0)
assert seen["flag"] == "--version"
Expand All @@ -110,6 +135,7 @@ def _fake(exe, flag, regex):
def test_get_rocm_version_parses_hipconfig_build_suffix(
monkeypatch, hipconfig_output, expected_version
):
monkeypatch.delenv("ROCM_VERSION", raising=False)
monkeypatch.setattr(C, "validateToolchain", lambda x: x)

class _R:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -1221,7 +1221,7 @@ def _capture_archs(*args, **kwargs):
monkeypatch.setattr(
RunModule,
"validateToolchain",
lambda *a: ("/fake/hipcc", None, "/fake/bundler", None, None),
lambda *a: ("/fake/hipcc", None, "/fake/bundler", None),
)
monkeypatch.setattr(RunModule, "makeIsaInfoMap", lambda _isas, _cxx: _stub_iim())
monkeypatch.setattr(RunModule, "assignGlobalParameters", _capture_gp)
Expand Down Expand Up @@ -1843,7 +1843,7 @@ def _capture_write(filename, *a, **kw):
monkeypatch.setattr(
RunModule,
"validateToolchain",
lambda *a: ("/fake/hipcc", None, "/fake/bundler", None, None),
lambda *a: ("/fake/hipcc", None, "/fake/bundler", None),
)
monkeypatch.setattr(RunModule, "makeIsaInfoMap", _iim)
monkeypatch.setattr(RunModule, "assignGlobalParameters", lambda *a, **kw: None)
Expand Down
29 changes: 23 additions & 6 deletions projects/hipblaslt/tensilelite/Tensile/Toolchain/Component.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,7 +54,7 @@ def _invoke(args: List[str], desc: str=""):
return out


def _getVersion(executable: str, versionFlag: str, regex: str) -> str:
def _getVersion(executable: str, versionFlag: str, regex: str) -> SemanticVersion:
"""Compute the version string of a toolchain component.

Args:
Expand All @@ -73,20 +73,37 @@ def _getVersion(executable: str, versionFlag: str, regex: str) -> str:
match = search(regex, output, IGNORECASE)
if match:
result = match.group(1)
return SemanticVersion(*[int(c.split("-")[0]) for c in result.split(".")[:3]])
return _parseVersion(result)
raise Exception(f"No version from {output} matches regex {regex}")
except Exception as e:
raise RuntimeError(f"Failed to get version when calling {args}: {e}")


def get_rocm_version() -> str:
"""Compute the ROCm version string using hipconfig.
def _parseVersion(version: str) -> SemanticVersion:
"""Parse the numeric major, minor, and patch prefix from a version string."""
version_match = search(r"^(\d+)\.(\d+)\.(\d+)", version.strip())
if not version_match:
raise ValueError(f"Invalid version string: {version}")
return SemanticVersion(*(int(component) for component in version_match.groups()))


def get_rocm_version() -> SemanticVersion:
"""Return the HIP package version supplied by CMake or reported by hipconfig.

Configured hipBLASLt builds pass CMake's ``hip_VERSION`` as
``ROCM_VERSION``. Standalone callers retain the existing ``hipconfig``
fallback.

Raises:
RuntimeError: If hipconfig fails to execute.
RuntimeError: If ROCM_VERSION is invalid or hipconfig fails to execute.
Return:
ROCm version string
ROCm SemanticVersion
"""
if version := environ.get("ROCM_VERSION"):
try:
return _parseVersion(version)
except ValueError as error:
raise RuntimeError(f"Invalid ROCM_VERSION: {version}") from error
return _getVersion(ToolchainDefaults.HIP_CONFIG, "--version", r'(.+)')


Expand Down
Loading