#!/bin/bash
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --cpus-per-task=8
#SBATCH --mem=50G
#SBATCH --gpus=A40:1
#SBATCH --time=2-00:00:00
#SBATCH --output=logs/%x_job_%j.out
#SBATCH --job-name=aefm_train
#SBATCH --partition=slowlane

module load Miniconda3
source ${EBROOTMINICONDA3}/bin/activate
conda activate aefm

time_date=$(date +%Y_%m_%d_%H_%M_%S)
echo "Job started on: ${time_date}"

mkdir -p ${SCRATCH_DIR}/temp_train
tmp_folder=${SCRATCH_DIR}/temp_train/$time_date

echo "Starting training"
echo "Temporary data folder at: $tmp_folder"

available_cpus=$(nproc)
echo "Available CPUs: $available_cpus"

# Detailed logging
export HYDRA_FULL_ERROR=1

# pretrain on t1x
# run_id=aefm_t1x
# aefm_train experiment=train_t1x \
#     run.id=${run_id} \
#     data.data_workdir=$tmp_folder \
#     data.num_workers=$available_cpus \

# finetune on t1x xtb
# run_id=aefm_t1x_xtb_finetune
# aefm_train experiment=train_t1x_xtb \
#     +pretrained=aefm/runs/aefm_t1x/best_model \
#     run.id=${run_id} \
#     data.data_workdir=$tmp_folder \
#     data.num_workers=$available_cpus \

# train from scratch on t1x xtb
# run_id=aefm_t1x_xtb
# aefm_train experiment=train_t1x_xtb \
#     run.id=${run_id} \
#     data.data_workdir=$tmp_folder \
#     data.num_workers=$available_cpus \

#! Fine-tuned
# run_id=aefm_t1x_with_swap_finetune
# aefm_train experiment=train_t1x_with_swap \
#     +pretrained=aefm/runs/aefm_t1x/best_model \
#     globals.sigma=0.24 \
#     run.id=${run_id} \
#     data.data_workdir=$tmp_folder \
#     data.num_workers=$available_cpus \

run_id=aefm_t1x_with_tmc_finetune
aefm_train experiment=train_t1x_with_tmc \
    +pretrained=aefm/runs/aefm_t1x/best_model \
    globals.sigma=0.24 \
    run.id=${run_id} \
    data.data_workdir=$tmp_folder \
    data.num_workers=$available_cpus \