diff --git a/docker/rocm.Dockerfile b/docker/rocm.Dockerfile index 6dcaa5af8..8d5e8c3a9 100644 --- a/docker/rocm.Dockerfile +++ b/docker/rocm.Dockerfile @@ -91,9 +91,6 @@ ARG BRANCH_TYPE=remote # Version override for setuptools_scm (used in nightly builds) ARG SETUPTOOLS_SCM_PRETEND_VERSION="" -ARG TRITON_REPO="https://github.com/triton-lang/triton.git" -ARG TRITON_COMMIT="42270451990532c67e69d753fbd026f28fcc4840" - ARG AITER_REPO="https://github.com/ROCm/aiter.git" ARG AITER_COMMIT="" ENV AITER_COMMIT="${AITER_COMMIT:-${AITER_COMMIT_DEFAULT}}" @@ -222,8 +219,8 @@ RUN if [ "$BUILD_LLVM" = "1" ]; then \ # leak into AITER's version when AITER uses setuptools_scm) ENV SETUPTOOLS_SCM_PRETEND_VERSION= -# Keep the base image's Torch-compatible Triton by default. Override with -# AITER_USE_SYSTEM_TRITON=0 when intentionally testing aiter-managed Triton. +# Compile AITER against the base image's Triton; the Triton step at the end of +# this file swaps in AITER's own pin afterwards. ENV AITER_USE_SYSTEM_TRITON=1 RUN pip uninstall -y aiter # Use `checkout -f` so the smudge-filter-induced "dirty" working tree from @@ -628,24 +625,6 @@ RUN cd /tmp/whl \ ;; \ esac - -# ----------------------- -# Hot patch: Triton -# For ROCm 7.2, this custom build breaks pip dependency management, -# so future `pip install` will break the ROCm stack. -# A workaround for this is to reinstall the default triton -# wheel with the `rocm/pytorch` image in the root directory. -RUN if [ "$BUILD_TRITON" = "1" ]; then \ - pip uninstall -y triton \ - && apt install -y cmake \ - && git clone ${TRITON_REPO} triton-custom \ - && cd triton-custom \ - && git checkout ${TRITON_COMMIT} \ - && pip install -r python/requirements.txt \ - && pip install -e . \ - && if [ -d python/triton_kernels ]; then pip install -e python/triton_kernels --no-deps; fi; \ - fi - # ----------------------- # Hot patch: transformers dynamic_module_utils symlink bug (v5.12.1). # _compute_local_source_files_hash calls Path(...).resolve() on custom-code @@ -676,6 +655,20 @@ else: print("patched transformers dynamic_module_utils.py (symlink hash fix)") PY +# ----------------------- +# Install the Triton AITER pins, replacing the base image's. No version check +# on purpose: the pin is AITER's to move, and its installer enforces a floor. +# +# Keep this last. Base ROCm Torch pins triton==3.5.1 and the torch patch above +# is what drops that pin, so installing Triton any earlier lets the next pip +# install pull CUDA torch instead. The hip check below is the tripwire. +RUN if [ "$BUILD_TRITON" = "1" ]; then \ + cd /sgl-workspace/aiter \ + && test -f .github/scripts/install_triton.sh \ + && PIP_NO_CACHE_DIR=1 bash .github/scripts/install_triton.sh \ + && python3 -c "import torch; from importlib.metadata import version; v = version('triton'); k = version('triton-kernels'); assert torch.version.hip is not None, torch.__version__; print(f'[Triton] ROCm Torch {torch.__version__}, Triton {v}, triton-kernels {k}')"; \ + fi + # ----------------------- # Performance environment variable.