Skip to content
Merged
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
123 changes: 90 additions & 33 deletions apps/rocm/src/therock.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,9 @@ use std::path::{Path, PathBuf};
use std::process::{Command, Output, Stdio};
use std::time::Duration;

const THEROCK_PIP_INDEX_BASE: &str = "https://rocm.nightlies.amd.com/v2";
const THEROCK_NIGHTLY_PIP_INDEX_BASE: &str = "https://rocm.nightlies.amd.com/v2";
const THEROCK_RELEASE_PIP_INDEX_BASE: &str = "https://repo.amd.com/rocm/whl";
const THEROCK_RELEASE_PIP_MULTI_ARCH_INDEX_BASE: &str = "https://repo.amd.com/rocm/whl-multi-arch";
const THEROCK_RELEASE_TARBALL_BASE: &str = "https://repo.amd.com/rocm/tarball/";
const THEROCK_NIGHTLY_TARBALL_BASE: &str = "https://rocm.nightlies.amd.com/tarball/";
const DEFAULT_MANAGED_PYTHON_VERSION: &str = "3.12";
Expand Down Expand Up @@ -1109,26 +1111,67 @@ fn resolve_pip_runtime_with_timeout(
download_timeout_secs: Option<u64>,
) -> Result<PipRuntimeResolution> {
let family_resolution = resolve_family(paths, family_override)?;
let index_url = therock_index_url(&family_resolution.family);
let index_urls = therock_index_urls(channel, &family_resolution.family);
let mut errors = Vec::new();
for index_url in index_urls {
match resolve_pip_runtime_from_index(
paths,
channel,
&family_resolution,
&index_url,
wheel_compatibility,
version_selector,
download_timeout_secs,
) {
Ok(resolution) => return Ok(resolution),
Err(error) => errors.push(format!("{index_url}: {error}")),
}
}
bail!(
"failed to resolve TheRock {} wheel runtime from candidate indexes:\n - {}",
channel.as_str(),
errors.join("\n - ")
)
}

fn resolve_pip_runtime_from_index(
paths: &AppPaths,
channel: TheRockChannel,
family_resolution: &FamilyResolution,
index_url: &str,
wheel_compatibility: &WheelCompatibility,
version_selector: Option<&RuntimeVersionSelector>,
download_timeout_secs: Option<u64>,
) -> Result<PipRuntimeResolution> {
let rocm_versions =
load_simple_index_versions(paths, &index_url, "rocm", None, download_timeout_secs)?;
load_simple_index_versions(paths, index_url, "rocm", None, download_timeout_secs)?;
if matches!(channel, TheRockChannel::Release)
&& version_selector.is_none()
&& !rocm_versions
.iter()
.any(|version| is_stable_runtime_version(version))
{
bail!(
"release channel only installs stable TheRock wheel versions, but no stable `rocm` package versions were found in {index_url}; try `rocm install sdk --channel release --format tarball` for stable release artifacts, or use `--channel nightly --format wheel` for preview builds"
);
}
let torch_versions = load_simple_index_versions(
paths,
&index_url,
index_url,
"torch",
Some(wheel_compatibility),
download_timeout_secs,
)?;
let torchvision_versions = load_simple_index_versions(
paths,
&index_url,
index_url,
"torchvision",
Some(wheel_compatibility),
download_timeout_secs,
)?;
let torchaudio_versions = load_simple_index_versions(
paths,
&index_url,
index_url,
"torchaudio",
Some(wheel_compatibility),
download_timeout_secs,
Expand All @@ -1149,9 +1192,9 @@ fn resolve_pip_runtime_with_timeout(
})?;
let latest_version = package_versions.rocm.clone();
Ok(PipRuntimeResolution {
family: family_resolution.family,
family_source: family_resolution.source,
index_url,
family: family_resolution.family.clone(),
family_source: family_resolution.source.clone(),
index_url: index_url.to_owned(),
latest_version,
package_versions,
})
Expand Down Expand Up @@ -1312,20 +1355,19 @@ fn channel_rocm_candidates(versions: &[String], channel: TheRockChannel) -> Vec<
let mut all = versions.to_vec();
all.sort_by(|left, right| compare_version_strings(left, right));
if matches!(channel, TheRockChannel::Release) {
let stable = all
return all
.iter()
.filter(|version| {
parse_version(version).is_some_and(|parsed| parsed.stage == VersionStage::Stable)
})
.filter(|version| is_stable_runtime_version(version))
.cloned()
.collect::<Vec<_>>();
if !stable.is_empty() {
return stable;
}
}
all
}

fn is_stable_runtime_version(version: &str) -> bool {
parse_version(version).is_some_and(|parsed| parsed.stage == VersionStage::Stable)
}

fn select_latest_stack_package(
versions: &[String],
rocm_version: &str,
Expand Down Expand Up @@ -1646,13 +1688,13 @@ fn select_latest_version(versions: &[String], channel: TheRockChannel) -> Option
let mut all = versions.to_vec();
all.sort_by(|left, right| compare_version_strings(left, right));
for version in versions {
if parse_version(version).is_some_and(|parsed| parsed.stage == VersionStage::Stable) {
if is_stable_runtime_version(version) {
stable.push(version.clone());
}
}
stable.sort_by(|left, right| compare_version_strings(left, right));
match channel {
TheRockChannel::Release => stable.pop().or_else(|| all.pop()),
TheRockChannel::Release => stable.pop(),
TheRockChannel::Nightly => all.pop(),
}
}
Expand Down Expand Up @@ -3263,8 +3305,14 @@ fn parse_version(value: &str) -> Option<ParsedVersion> {
})
}

fn therock_index_url(family: &str) -> String {
format!("{THEROCK_PIP_INDEX_BASE}/{family}")
fn therock_index_urls(channel: TheRockChannel, family: &str) -> Vec<String> {
match channel {
TheRockChannel::Release => vec![
format!("{THEROCK_RELEASE_PIP_INDEX_BASE}/{family}"),
format!("{THEROCK_RELEASE_PIP_MULTI_ARCH_INDEX_BASE}/{family}"),
],
TheRockChannel::Nightly => vec![format!("{THEROCK_NIGHTLY_PIP_INDEX_BASE}/{family}")],
}
}

const fn platform_tarball_token() -> &'static str {
Expand Down Expand Up @@ -3420,6 +3468,15 @@ mod tests {
);
}

#[test]
fn release_channel_rejects_prerelease_only_versions() {
let versions = vec!["7.13.0a20260326".to_owned(), "7.14.0rc1".to_owned()];
assert_eq!(
select_latest_version(&versions, TheRockChannel::Release),
None
);
}

#[test]
fn pip_runtime_installs_pinned_devel_and_torch_stack_from_therock_index() {
let package_versions = TheRockPipPackageVersions {
Expand All @@ -3445,21 +3502,21 @@ mod tests {
#[test]
fn pip_runtime_selects_latest_common_rocm_suffix_not_latest_rocm_package() {
let rocm_versions = vec![
"7.13.0a20260512".to_owned(),
"7.13.0a20260513".to_owned(),
"7.14.0a20260602".to_owned(),
"7.13.0".to_owned(),
"7.13.1".to_owned(),
"7.14.0".to_owned(),
];
let torch_versions = vec![
"2.9.1+rocm7.13.0a20260513".to_owned(),
"2.10.0+rocm7.13.0a20260513".to_owned(),
"2.9.1+rocm7.13.1".to_owned(),
"2.10.0+rocm7.13.1".to_owned(),
];
let torchvision_versions = vec![
"0.24.0+rocm7.13.0a20260513".to_owned(),
"0.25.0+rocm7.13.0a20260513".to_owned(),
"0.24.0+rocm7.13.1".to_owned(),
"0.25.0+rocm7.13.1".to_owned(),
];
let torchaudio_versions = vec![
"2.9.0+rocm7.13.0a20260513".to_owned(),
"2.10.0+rocm7.13.0a20260513".to_owned(),
"2.9.0+rocm7.13.1".to_owned(),
"2.10.0+rocm7.13.1".to_owned(),
];

let selected = select_matching_pip_package_versions(
Expand All @@ -3472,10 +3529,10 @@ mod tests {
)
.expect("expected compatible package set");

assert_eq!(selected.rocm, "7.13.0a20260513");
assert_eq!(selected.torch, "2.10.0+rocm7.13.0a20260513");
assert_eq!(selected.torchvision, "0.25.0+rocm7.13.0a20260513");
assert_eq!(selected.torchaudio, "2.10.0+rocm7.13.0a20260513");
assert_eq!(selected.rocm, "7.13.1");
assert_eq!(selected.torch, "2.10.0+rocm7.13.1");
assert_eq!(selected.torchvision, "0.25.0+rocm7.13.1");
assert_eq!(selected.torchaudio, "2.10.0+rocm7.13.1");
}

#[test]
Expand Down
Loading
Loading