ARG UBUNTU_VERSION=22.04

FROM ubuntu:${UBUNTU_VERSION} AS base

ARG UBUNTU_VERSION

ENV DEBIAN_FRONTEND=noninteractive

ARG PIP_EXTRA_INDEX_URL
ENV PIP_EXTRA_INDEX_URL=${PIP_EXTRA_INDEX_URL}
ARG PIP_PREFER_BINARY
ENV PIP_PREFER_BINARY=${PIP_PREFER_BINARY}

ARG CLANG_VERSION

# Allow APT to run as root
RUN echo 'APT::Sandbox::User "root";' | tee -a /etc/apt/apt.conf.d/10sandbox

# Install common dependencies (so that this step can be cached separately)
COPY ./common/install_base.sh install_base.sh
RUN bash ./install_base.sh && rm install_base.sh

# Install clang
ARG LLVMDEV
COPY ./common/install_clang.sh install_clang.sh
RUN bash ./install_clang.sh && rm install_clang.sh

# Install user
COPY ./common/install_user.sh install_user.sh
RUN bash ./install_user.sh && rm install_user.sh

# Install katex
ARG KATEX
COPY ./common/install_docs_reqs.sh install_docs_reqs.sh
RUN bash ./install_docs_reqs.sh && rm install_docs_reqs.sh

# Install python (via deadsnakes) into a venv and the CI requirements
ARG PYTHON_VERSION
ARG PYTHON_FREETHREADED
ARG DOCS
ENV PYTHON_VERSION=$PYTHON_VERSION
ENV PYTHON_FREETHREADED=$PYTHON_FREETHREADED
ENV DOCS=$DOCS
ENV VENV_PATH=/opt/python-${PYTHON_VERSION}-venv
ENV VIRTUAL_ENV=${VENV_PATH}
ENV PATH=${VENV_PATH}/bin:$PATH
COPY requirements-ci.txt requirements-docs.txt /opt/
COPY ./common/install_python.sh install_python.sh
RUN bash ./install_python.sh && rm install_python.sh /opt/requirements-ci.txt /opt/requirements-docs.txt

# Install gcc
ARG GCC_VERSION
COPY ./common/install_gcc.sh install_gcc.sh
RUN bash ./install_gcc.sh && rm install_gcc.sh

# Install lcov for C++ code coverage
COPY ./common/install_lcov.sh install_lcov.sh
RUN  bash ./install_lcov.sh && rm install_lcov.sh

# Install cuda and cudnn
ARG CUDA_VERSION
COPY ./common/install_cuda.sh install_cuda.sh
COPY ./common/install_nccl.sh install_nccl.sh
COPY ./ci_commit_pins/nccl* /ci_commit_pins/
COPY ./common/install_cusparselt.sh install_cusparselt.sh
RUN bash ./install_cuda.sh ${CUDA_VERSION} && rm install_cuda.sh install_nccl.sh /ci_commit_pins/nccl* install_cusparselt.sh
ENV DESIRED_CUDA=${CUDA_VERSION}
ENV PATH=/usr/local/nvidia/bin:/usr/local/cuda/bin:$PATH

# Install MAGMA into the CUDA toolkit dir (was previously pulled in via conda)
COPY ./common/install_magma.sh install_magma.sh
RUN if [ -n "${CUDA_VERSION}" ]; then bash ./install_magma.sh $(echo ${CUDA_VERSION} | cut -f1-2 -d'.'); fi
RUN rm -f install_magma.sh
ENV MAGMA_HOME=/usr/local/cuda/magma
# No effect if cuda not installed
ENV USE_SYSTEM_NCCL=1
ENV NCCL_INCLUDE_DIR="/usr/local/cuda/include/"
ENV NCCL_LIB_DIR="/usr/local/cuda/lib64/"


ARG CUDA_VERSION

ARG INDUCTOR_BENCHMARKS
COPY ./common/install_inductor_benchmark_deps.sh install_inductor_benchmark_deps.sh
COPY ./common/common_utils.sh common_utils.sh
COPY ci_commit_pins/huggingface-requirements.txt huggingface-requirements.txt
COPY ci_commit_pins/timm.txt timm.txt
COPY ci_commit_pins/torchbench.txt torchbench.txt
# Only build aoti cpp tests when INDUCTOR_BENCHMARKS is set to True
ENV BUILD_AOT_INDUCTOR_TEST=${INDUCTOR_BENCHMARKS}
RUN if [ -n "${INDUCTOR_BENCHMARKS}" ]; then bash ./install_inductor_benchmark_deps.sh; fi
RUN rm install_inductor_benchmark_deps.sh common_utils.sh timm.txt huggingface-requirements.txt torchbench.txt

ARG INSTALL_MINGW
COPY ./common/install_mingw.sh install_mingw.sh
RUN if [ -n "${INSTALL_MINGW}" ]; then bash ./install_mingw.sh; fi
RUN rm install_mingw.sh

ARG TRITON
ARG TRITON_CPU

# Create a separate stage for building Triton and Triton-CPU.  install_triton
# will check for the presence of env vars
FROM base AS triton-builder
COPY ./common/install_triton.sh install_triton.sh
COPY ./common/common_utils.sh common_utils.sh
COPY ci_commit_pins/triton.txt triton.txt
COPY ci_commit_pins/triton-cpu.txt triton-cpu.txt
RUN bash ./install_triton.sh

FROM base AS final
COPY --from=triton-builder /opt/triton /opt/triton
RUN if [ -n "${TRITON}" ] || [ -n "${TRITON_CPU}" ]; then pip install /opt/triton/*.whl; chown -R jenkins:jenkins ${VENV_PATH}; fi
RUN rm -rf /opt/triton

ARG EXECUTORCH
# Build and install executorch
COPY ./common/install_executorch.sh install_executorch.sh
COPY ./common/common_utils.sh common_utils.sh
COPY ci_commit_pins/executorch.txt executorch.txt
RUN if [ -n "${EXECUTORCH}" ]; then bash ./install_executorch.sh; fi
RUN rm install_executorch.sh common_utils.sh executorch.txt

ARG HALIDE
# Build and install halide
COPY ./common/install_halide.sh install_halide.sh
COPY ./common/common_utils.sh common_utils.sh
COPY ci_commit_pins/halide.txt halide.txt
RUN if [ -n "${HALIDE}" ]; then bash ./install_halide.sh; fi
RUN rm install_halide.sh common_utils.sh halide.txt

ARG PALLAS
ARG TPU
ARG CUDA_VERSION
# Install JAX (for Pallas) - with TPU/CUDA support based on args
# TPU=yes -> install JAX with TPU support
# PALLAS=yes + CUDA_VERSION -> install JAX with CUDA support
# PALLAS=yes (no CUDA) -> install JAX CPU-only
COPY ./common/requirements_tpu.txt requirements_tpu.txt
COPY ./common/install_jax.sh install_jax.sh
COPY ./common/common_utils.sh common_utils.sh
COPY ./ci_commit_pins/jax.txt /ci_commit_pins/jax.txt
RUN if [ -n "${TPU}" ]; then \
      bash -c "source ./common_utils.sh && pip_install -r requirements_tpu.txt" && \
      bash ./install_jax.sh tpu; \
    elif [ -n "${PALLAS}" ]; then \
      bash ./install_jax.sh ${CUDA_VERSION:-cpu}; \
    fi
RUN rm -f install_jax.sh common_utils.sh /ci_commit_pins/jax.txt

ARG ONNX
# Install ONNX dependencies
COPY ./common/install_onnx.sh ./common/common_utils.sh ./
RUN if [ -n "${ONNX}" ]; then bash ./install_onnx.sh; fi
RUN rm install_onnx.sh common_utils.sh

ARG TVM
# Install apache-tvm for the dynamo tvm backend tests
# apache-tvm declares apache-tvm-ffi>=0.1.12 with no upper bound, but its
# prebuilt libtvm_runtime.so is ABI-bound to the tvm-ffi it was built against,
# so the floating dep has to be pinned alongside it.
COPY ./common/common_utils.sh common_utils.sh
COPY ci_commit_pins/tvm.txt tvm.txt
COPY ci_commit_pins/tvm-ffi.txt tvm-ffi.txt
RUN if [ -n "${TVM}" ]; then bash -c 'source ./common_utils.sh && pip_install apache-tvm==$(get_pinned_commit tvm) apache-tvm-ffi==$(get_pinned_commit tvm-ffi) && python -c "import tvm"'; fi
RUN rm common_utils.sh tvm.txt tvm-ffi.txt

# Build TSan-instrumented CPython for thread sanitizer testing
ARG TSAN
ARG CLANG_VERSION
COPY ./common/install_cpython.sh install_cpython.sh
COPY requirements-ci.txt /tmp/requirements-ci.txt
RUN if [ -n "${TSAN}" ]; then \
      CC=clang-${CLANG_VERSION} CXX=clang++-${CLANG_VERSION} \
      CPYTHON_VERSIONS="3.14.4t+tsan" bash ./install_cpython.sh && \
      /opt/python/cp314-cp314t+tsan/bin/pip install -r /tmp/requirements-ci.txt; \
    fi
RUN rm -f install_cpython.sh /tmp/requirements-ci.txt

# (optional) Build ACL
ARG ACL
COPY ./common/install_acl.sh install_acl.sh
RUN if [ -n "${ACL}" ]; then bash ./install_acl.sh; fi
RUN rm install_acl.sh
ENV INSTALLED_ACL=${ACL}

ARG OPENBLAS
COPY ./common/install_openblas.sh install_openblas.sh
RUN if [ -n "${OPENBLAS}" ]; then bash ./install_openblas.sh; fi
RUN rm install_openblas.sh
ENV INSTALLED_OPENBLAS=${OPENBLAS}

# Install Rust toolchain (needed to build the torch._rust extension).
COPY ./common/install_rust.sh install_rust.sh
COPY ./ci_commit_pins/rust.txt rust.txt
ENV RUSTUP_HOME /opt/rust
ENV CARGO_HOME /opt/rust
ENV PATH /opt/rust/bin:$PATH
RUN bash ./install_rust.sh && rm install_rust.sh rust.txt

# Install ccache/sccache (do this last, so we get priority in PATH)
ARG SKIP_SCCACHE_INSTALL
COPY ./common/install_cache.sh install_cache.sh
COPY ./common/patches /opt/cache/patches
ENV PATH=/opt/cache/bin:$PATH
RUN if [ -z "${SKIP_SCCACHE_INSTALL}" ]; then bash ./install_cache.sh; fi
RUN rm install_cache.sh

# Add jni.h for java host build
COPY ./common/install_jni.sh install_jni.sh
COPY ./java/jni.h jni.h
RUN bash ./install_jni.sh && rm install_jni.sh

# Install Open MPI for CUDA
COPY ./common/install_openmpi.sh install_openmpi.sh
RUN if [ -n "${CUDA_VERSION}" ]; then bash install_openmpi.sh; fi
RUN rm install_openmpi.sh

# Include BUILD_ENVIRONMENT environment variable in image
ARG BUILD_ENVIRONMENT
ENV BUILD_ENVIRONMENT=${BUILD_ENVIRONMENT}

# AWS specific CUDA build guidance
ENV TORCH_NVCC_FLAGS="-Xfatbin -compress-all"
ENV CUDA_PATH=/usr/local/cuda

USER jenkins
CMD ["bash"]
