FROM nvidia/cuda:12.2.0-devel-ubuntu20.04

RUN apt-get update && apt-get install --yes \
      python3 \
      python3-distutils \
      python3-pip \
      clang \
      wget \
      vim \
      git

RUN python3 -m pip install --ignore-installed \
      "clang~=$(clang --version | grep -oP '10\.[^-]+')" \
      torch \
      torchvision \
      lightning \
      numpy \
      memory_profiler

ENV PYTORCH_DATASETS_DIR=/pytorch-data
ENV TORCH_HOME=/pytorch-home
COPY download_pytorch_datasets.py /tmp/
# Some PyTorch examples hardcode the data directory to "data", so
# make a symlink for that too.
RUN mkdir "$PYTORCH_DATASETS_DIR" && \
    python3 /tmp/download_pytorch_datasets.py && \
    rm /tmp/download_pytorch_datasets.py

RUN PYTORCH_EXAMPLES_COMMIT=30b310a977a82dbfc3d8e4a820f3b14d876d3bd2 && \
    mkdir /pytorch-examples && \
    cd /pytorch-examples && \
    git init && \
    git remote add origin https://github.com/pytorch/examples && \
    git fetch --depth 1 origin "$PYTORCH_EXAMPLES_COMMIT" && \
    git checkout FETCH_HEAD && \
    sed -ri "s~(datasets.*)\\(['\"](../)?data['\"],~\\1('$PYTORCH_DATASETS_DIR',~g" **/*.py && \
    sed -ri 's/download=True/download=False/' **/*.py

COPY *.py /
RUN rm /download_pytorch_datasets.py
