Views
No views yet
torch state_dicts. The subdirectories arc_v1_public and arc_v2_public contain the final checkpoints step_<final-step>, which can be loaded with the load_checkpoint or by providing the checkpoint path as load_checkpoint=path/to/checkpoint. For reference, see the PretrainConfig in pretrain.py.1# use uv for venv
2sudo snap install astral-uv --classic
3uv venv .venv -p 3.12
4source .venv/bin/activate
5
6# install python-dev for adam atan2
7sudo apt install python3-dev -y
8# install torch
9PYTORCH_INDEX_URL=https://download.pytorch.org/whl/cu128
10uv pip install torch torchvision torchaudio --index-url $PYTORCH_INDEX_URL
11# install dependencies + adam atan
12uv pip install packaging ninja wheel setuptools setuptools-scm
13uv pip install --no-cache-dir --no-build-isolation adam-atan2
14
15# test torch, cuda and AdamAtan2
16python
17import torch
18t = torch.tensor([0,1,2]).to('cuda')
19from adam_atan2 import AdamATan2
20
21# install remaining dependencies
22uv pip install -r requirements.txt1python -m dataset.build_arc_dataset \
2 --input-file-prefix kaggle/combined/arc-agi \
3 --output-dir data/arc1concept-aug-1000 \
4 --subsets training evaluation concept \
5 --test-set-name evaluation1python -m dataset.build_arc_dataset \
2 --input-file-prefix kaggle/combined/arc-agi \
3 --output-dir data/arc2concept-aug-1000 \
4 --subsets training2 evaluation2 concept \
5 --test-set-name evaluation21run_name="trm_arc_v1_public"
2torchrun --nproc-per-node 8 --rdzv_backend=c10d --rdzv_endpoint=localhost:0 --nnodes=1 pretrain.py \
3arch=trm \
4data_paths="[data/arc1concept-aug-1000]" \
5arch.L_layers=2 \
6arch.H_cycles=3 arch.L_cycles=4 \
7+run_name=${run_name} ema=True1run_name="trm_arc_v2_public"
2torchrun --nproc-per-node 8 --rdzv_backend=c10d --rdzv_endpoint=localhost:0 --nnodes=1 pretrain.py \
3arch=trm \
4data_paths="[data/arc2concept-aug-1000]" \
5arch.L_layers=2 \
6arch.H_cycles=3 arch.L_cycles=4 \
7+run_name=${run_name} ema=True1export MAIN_ADDR=<MAIN_IP>
2export MAIN_PORT=29500
3export NNODES=2
4export GPUS_PER_NODE=8
5export OMP_NUM_THREADS=8
6export NCCL_PORT_RANGE=40000-40050
7run_name="arc_v1_public_2_nodes"
8# on each node:
9export NODE_RANK=0
10torchrun \
11 --nnodes $NNODES \
12 --node_rank $NODE_RANK \
13 --nproc_per_node $GPUS_PER_NODE \
14 --rdzv_backend c10d \
15 --rdzv_endpoint $MAIN_ADDR:$MAIN_PORT \
16 pretrain.py \
17arch=trm \
18data_paths="[data/arc1concept-aug-1000]" \
19arch.L_layers=2 \
20arch.H_cycles=3 arch.L_cycles=4 \
21+run_name=${run_name} ema=True \
22eval_interval=50000