diff --git a/projects/hipblaslt/cmake/hipblaslt_python.cmake b/projects/hipblaslt/cmake/hipblaslt_python.cmake index 658f50fe23cb..0eb4c71e01fe 100644 --- a/projects/hipblaslt/cmake/hipblaslt_python.cmake +++ b/projects/hipblaslt/cmake/hipblaslt_python.cmake @@ -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}" diff --git a/projects/hipblaslt/tensilelite/Tensile/TensileCreateLibrary/Run.py b/projects/hipblaslt/tensilelite/Tensile/TensileCreateLibrary/Run.py index 9cafd1dc4586..22d9678404a8 100644 --- a/projects/hipblaslt/tensilelite/Tensile/TensileCreateLibrary/Run.py +++ b/projects/hipblaslt/tensilelite/Tensile/TensileCreateLibrary/Run.py @@ -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 @@ -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"]: diff --git a/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/TensileCreateLibraryRun/test_r7_createlib_deep_char.py b/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/TensileCreateLibraryRun/test_r7_createlib_deep_char.py index edd00cfabf77..ae1af753a2a7 100644 --- a/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/TensileCreateLibraryRun/test_r7_createlib_deep_char.py +++ b/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/TensileCreateLibraryRun/test_r7_createlib_deep_char.py @@ -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()), \ @@ -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) @@ -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()), \ @@ -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()), \ @@ -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()), \ @@ -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()), \ @@ -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()), \ @@ -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()), \ @@ -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()), \ diff --git a/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/ToolchainComponent/test_toolchain_component_char.py b/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/ToolchainComponent/test_toolchain_component_char.py index 560a3347e7de..affdaaff40a1 100644 --- a/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/ToolchainComponent/test_toolchain_component_char.py +++ b/projects/hipblaslt/tensilelite/Tensile/Tests/unit/characterization/ToolchainComponent/test_toolchain_component_char.py @@ -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" @@ -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: diff --git a/projects/hipblaslt/tensilelite/Tensile/Tests/unit/test_gfx1250_asic_revision.py b/projects/hipblaslt/tensilelite/Tensile/Tests/unit/test_gfx1250_asic_revision.py index bf533feafcbb..a11f73cdccea 100644 --- a/projects/hipblaslt/tensilelite/Tensile/Tests/unit/test_gfx1250_asic_revision.py +++ b/projects/hipblaslt/tensilelite/Tensile/Tests/unit/test_gfx1250_asic_revision.py @@ -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) @@ -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) diff --git a/projects/hipblaslt/tensilelite/Tensile/Toolchain/Component.py b/projects/hipblaslt/tensilelite/Tensile/Toolchain/Component.py index fd5a986f444a..90f47589169f 100644 --- a/projects/hipblaslt/tensilelite/Tensile/Toolchain/Component.py +++ b/projects/hipblaslt/tensilelite/Tensile/Toolchain/Component.py @@ -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: @@ -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'(.+)')