-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathDockerfile.train
More file actions
48 lines (36 loc) · 2.08 KB
/
Copy pathDockerfile.train
File metadata and controls
48 lines (36 loc) · 2.08 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
# syntax=docker/dockerfile:1
# A plain Python base is all we need: torch and its CUDA 12.8 libraries are installed from uv.lock by `uv sync`,
# so only the host driver matters. cu128 wheels cover Turing (sm_75) through Blackwell (sm_120) and run on host
# drivers >= R525 -- but they dropped Volta, so avoid V100/Pascal hosts on vast.ai.
#
# IMPORTANT: build with `--platform linux/amd64` (vast.ai hosts are x86_64). The lockfile only resolves to
# CUDA-enabled torch on x86_64 linux -- an arm64 build (e.g. on an Apple Silicon Mac without --platform) silently
# gets CPU-only torch.
FROM python:3.12-slim-bookworm
ENV PYTHONUNBUFFERED=1
ENV UV_LINK_MODE=copy
ENV UV_COMPILE_BYTECODE=1
ENV UV_PROJECT_ENVIRONMENT=/.venv
COPY --from=ghcr.io/astral-sh/uv:latest /uv /uvx /bin/
WORKDIR /
COPY pyproject.toml uv.lock /
RUN --mount=type=cache,target=/root/.cache/uv \
uv sync --frozen
ENV PATH="/.venv/bin:$PATH"
# `spacy download` shells out to pip, which uv-created venvs don't include; install the model wheel directly
RUN uv pip install --python /.venv/bin/python \
https://github.com/explosion/spacy-models/releases/download/en_core_web_md-3.8.0/en_core_web_md-3.8.0-py3-none-any.whl
# Pre-bake Hugging Face download (base model) so fresh instances don't refetch them. Only the default --model_base
ENV HF_HOME=/hf
RUN python -c "from transformers import AutoModelForSequenceClassification, AutoTokenizer; \
AutoTokenizer.from_pretrained('bert-base-uncased'); \
AutoModelForSequenceClassification.from_pretrained('bert-base-uncased', num_labels=2)"
COPY . /.
ENV PYTHONPATH=/
# NOTE: --train_final expects data/models/best/{case_id}/thresholds.json (produced by thresholds.py from a prior
# CV run). It is NOT baked into this image -- copy it onto the instance before final training.
# Runtime env vars to set on the instance: WANDB_API_KEY, plus AWS credentials for S3 uploads / SQS parallel mode.
# Using ENTRYPOINT allows us to pass args with `docker run`
# ENTRYPOINT ["python", "/src/sent_spans/train.py"]
# For now we'll manually start training, so just sleep here
CMD ["sleep", "infinity"]