Skip to content

Commit 93bbd3a

Browse files
committed
Ship a single ROCm version to the nightly getting-started matrix
The getting-started page renders one ROCm choice. gen_quick_start_module.py picks it with max() over version strings, and "7.14" < "7.2" lexicographically, so sending both made pytorch.org advertise ROCm 7.2 as the nightly ROCm after 7.14 landed. Send only the newest, chosen with a version-aware key. Scoped to getting-started + nightly: the binary build matrix is untouched, so nightly still builds every ROCm in ROCM_ARCHES_DICT.
1 parent 00fa578 commit 93bbd3a

2 files changed

Lines changed: 51 additions & 1 deletion

File tree

tools/scripts/generate_binary_build_matrix.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -131,6 +131,10 @@
131131
DISABLE = "disable"
132132

133133

134+
def parse_version(version: str) -> Tuple[int, ...]:
135+
return tuple(int(part) for part in version.split("."))
136+
137+
134138
def arch_type(arch_version: str) -> str:
135139
if arch_version in CUDA_ARCHES:
136140
return CUDA
@@ -182,6 +186,8 @@ def initialize_globals(
182186

183187
CUDA_ARCHES = CUDA_ARCHES_DICT[channel]
184188
ROCM_ARCHES = ROCM_ARCHES_DICT[channel]
189+
if getting_started and channel == NIGHTLY:
190+
ROCM_ARCHES = [max(ROCM_ARCHES, key=parse_version)]
185191
if build_python_only:
186192
# Only select the oldest version of python if building a python only package
187193
PYTHON_ARCHES = [PYTHON_ARCHES_DICT[channel][0]]

tools/tests/test_generate_binary_build_matrix.py

Lines changed: 45 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,11 @@
44
import sys
55
from unittest import main, TestCase
66

7-
from tools.scripts.generate_binary_build_matrix import generate_build_matrix
7+
from tools.scripts.generate_binary_build_matrix import (
8+
generate_build_matrix,
9+
parse_version,
10+
ROCM_ARCHES_DICT,
11+
)
812

913

1014
ASSETS_DIR = "tools/tests/assets"
@@ -200,6 +204,46 @@ def test_torch_only_install_command_for_torch_only_arches(self):
200204
self.assertNotIn("torchvision", entry["installation"])
201205
self.assertIn("torch", entry["installation"])
202206

207+
def _rocm_versions(self, channel: str, getting_started: str) -> set:
208+
out = generate_build_matrix(
209+
"wheel",
210+
"linux",
211+
channel,
212+
"enable",
213+
"enable",
214+
"enable",
215+
"enable",
216+
"false",
217+
"false",
218+
"disable",
219+
getting_started,
220+
None,
221+
)
222+
return {
223+
entry["gpu_arch_version"]
224+
for entry in out["include"]
225+
if entry["gpu_arch_type"] == "rocm"
226+
}
227+
228+
def test_parse_version_orders_double_digit_minors(self):
229+
self.assertGreater(parse_version("7.14"), parse_version("7.2"))
230+
self.assertEqual(
231+
max(["7.2", "7.14"], key=parse_version),
232+
"7.14",
233+
)
234+
235+
def test_getting_started_nightly_ships_one_rocm(self):
236+
versions = self._rocm_versions("nightly", "true")
237+
self.assertEqual(len(versions), 1)
238+
self.assertEqual(
239+
versions, {max(ROCM_ARCHES_DICT["nightly"], key=parse_version)}
240+
)
241+
242+
def test_nightly_builds_keep_every_rocm(self):
243+
versions = self._rocm_versions("nightly", "false")
244+
self.assertEqual(versions, set(ROCM_ARCHES_DICT["nightly"]))
245+
self.assertGreater(len(versions), 1)
246+
203247

204248
def parse_args():
205249
parser = argparse.ArgumentParser(description="Test generate build matrix")

0 commit comments

Comments
 (0)