Initial commit: DoVLA-CIL codebase (h=16 breakthrough) (part 2)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- scripts/slurm/generate_6task_h16.sbatch +87 -0
- scripts/slurm/generate_cil_array.sbatch +63 -0
- scripts/slurm/generate_embeddings.sbatch +28 -0
- scripts/slurm/horizon_sweep_pickcube.sbatch +88 -0
- scripts/slurm/install_smolvla_env.sbatch +104 -0
- scripts/slurm/make_maniskill_collection.sbatch +32 -0
- scripts/slurm/maniskill_lattice_debug.sbatch +111 -0
- scripts/slurm/maniskill_lattice_full.sbatch +80 -0
- scripts/slurm/maniskill_multitask_pilot.sbatch +63 -0
- scripts/slurm/phase_a1_generate_10k.sbatch +93 -0
- scripts/slurm/phase_a1_generate_10k_enhanced.sbatch +120 -0
- scripts/slurm/phase_a1_revised_enhanced.sbatch +65 -0
- scripts/slurm/phase_a1b_train_enhanced.sbatch +63 -0
- scripts/slurm/phase_a2_train_large_model.sbatch +59 -0
- scripts/slurm/phase_a3_eval_large_model.sbatch +50 -0
- scripts/slurm/phase_a4_hparam_sweep.sbatch +65 -0
- scripts/slurm/phase_a5_horizon_sweep.sbatch +63 -0
- scripts/slurm/phase_b_generate_12tasks.sbatch +109 -0
- scripts/slurm/phase_b_train_12tasks.sbatch +64 -0
- scripts/slurm/plan_c_generate_10k.sbatch +121 -0
- scripts/slurm/prepare_maniskill_baselines.sbatch +33 -0
- scripts/slurm/render_maniskill_multitask.sbatch +33 -0
- scripts/slurm/render_maniskill_observations.sbatch +49 -0
- scripts/slurm/run_external_vla_baseline.sbatch +51 -0
- scripts/slurm/run_scaling.sbatch +42 -0
- scripts/slurm/run_smolvla_cil_baseline.sbatch +48 -0
- scripts/slurm/smoke_smolvla_checkpoint.sbatch +42 -0
- scripts/slurm/train_attention_model.sbatch +62 -0
- scripts/slurm/train_dovla.sbatch +44 -0
- scripts/slurm/train_enhanced_model.sbatch +63 -0
- scripts/slurm/train_h16_policy.sbatch +54 -0
- scripts/slurm/train_hybrid_direct.sbatch +65 -0
- scripts/slurm/train_maniskill_baseline_array.sbatch +84 -0
- scripts/slurm/train_maniskill_collection_array.sbatch +114 -0
- scripts/slurm/train_maniskill_collection_cpu_array.sbatch +61 -0
- scripts/slurm/train_maniskill_debug.sbatch +86 -0
- scripts/slurm/train_maniskill_full_array.sbatch +74 -0
- scripts/slurm/train_maniskill_scaling_array.sbatch +82 -0
- scripts/slurm/train_maniskill_visual_array.sbatch +79 -0
- scripts/slurm/train_transformer.sbatch +72 -0
- scripts/slurm/train_transformer_lang.sbatch +68 -0
- scripts/smoke_full_pipeline.py +155 -0
- scripts/smoke_smolvla_checkpoint.py +169 -0
- scripts/smoke_test.sh +27 -0
- scripts/train_dovla.py +153 -0
- scripts/train_dovla_attention.py +324 -0
- scripts/train_dovla_enhanced.py +407 -0
- scripts/train_dovla_transformer.py +366 -0
- scripts/train_hybrid_direct.py +348 -0
- scripts/train_transformer_with_language.py +368 -0
scripts/slurm/generate_6task_h16.sbatch
ADDED
|
@@ -0,0 +1,87 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_6task_h16
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=08:00:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 12 |
+
#SBATCH --array=0-5
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
# Generate 6-task CIL collection with horizon=16 (vs baseline h=4)
|
| 17 |
+
# Expected: oracle ceiling ~90%+ (vs 42.57% @ h=4)
|
| 18 |
+
# This enables policy success 50-70%+ (vs 29.67% @ h=4)
|
| 19 |
+
|
| 20 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 21 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 22 |
+
SIF="$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif"
|
| 23 |
+
PYTHON="$SCRATCH_ROOT/envs/maniskill/bin/python"
|
| 24 |
+
NATIVE_LIBS="$SCRATCH_ROOT/native_libs/lib"
|
| 25 |
+
CPU_RENDER_LIBS="$SCRATCH_ROOT/cpu_render_libs"
|
| 26 |
+
CA_BUNDLE="$SCRATCH_ROOT/ca-bundle.crt"
|
| 27 |
+
VULKAN_ICD="$CPU_RENDER_LIBS/share/vulkan/icd.d/lvp_icd.x86_64.json"
|
| 28 |
+
OUT_ROOT="${OUT_ROOT:-$SCRATCH_ROOT/experiments/six_task_h16_collection}"
|
| 29 |
+
RUNTIME_DIR="/tmp/$USER/dovla-runtime-$SLURM_JOB_ID"
|
| 30 |
+
CACHE_DIR="/tmp/$USER/dovla-mesa-$SLURM_JOB_ID"
|
| 31 |
+
|
| 32 |
+
# Task array
|
| 33 |
+
TASKS=(PickCube-v1 PushCube-v1 PullCube-v1 StackCube-v1 LiftPegUpright-v1 PegInsertionSide-v1)
|
| 34 |
+
TASK=${TASKS[$SLURM_ARRAY_TASK_ID]}
|
| 35 |
+
|
| 36 |
+
# Demo paths
|
| 37 |
+
declare -A DEMOS
|
| 38 |
+
DEMOS[PickCube-v1]="$SCRATCH_ROOT/maniskill_data/demos/PickCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 39 |
+
DEMOS[PushCube-v1]="$SCRATCH_ROOT/maniskill_multitask_demos/PushCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 40 |
+
DEMOS[PullCube-v1]="$SCRATCH_ROOT/maniskill_multitask_demos/PullCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 41 |
+
DEMOS[StackCube-v1]="$SCRATCH_ROOT/maniskill_multitask_demos/StackCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 42 |
+
DEMOS[LiftPegUpright-v1]="$SCRATCH_ROOT/maniskill_multitask_demos/LiftPegUpright-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 43 |
+
DEMOS[PegInsertionSide-v1]="$SCRATCH_ROOT/maniskill_multitask_demos/PegInsertionSide-v1/rl/trajectory.h5"
|
| 44 |
+
|
| 45 |
+
DEMO_PATH="${DEMOS[$TASK]}"
|
| 46 |
+
OUT_DIR="$OUT_ROOT/$TASK"
|
| 47 |
+
|
| 48 |
+
# Groups per task
|
| 49 |
+
if [[ "$TASK" == "PickCube-v1" ]]; then
|
| 50 |
+
NUM_GROUPS=1000
|
| 51 |
+
else
|
| 52 |
+
NUM_GROUPS=500
|
| 53 |
+
fi
|
| 54 |
+
|
| 55 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 56 |
+
cd "$PROJECT_DIR"
|
| 57 |
+
mkdir -p outputs/hpc/logs "$OUT_DIR" "$RUNTIME_DIR" "$CACHE_DIR"
|
| 58 |
+
chmod 700 "$RUNTIME_DIR"
|
| 59 |
+
|
| 60 |
+
export OMP_NUM_THREADS=1 OPENBLAS_NUM_THREADS=1 MKL_NUM_THREADS=1 LP_NUM_THREADS=1
|
| 61 |
+
|
| 62 |
+
ENVS="LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS:/.singularity.d/libs,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=1,SSL_CERT_FILE=$CA_BUNDLE,REQUESTS_CA_BUNDLE=$CA_BUNDLE,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1"
|
| 63 |
+
|
| 64 |
+
echo "=================================================="
|
| 65 |
+
echo "Task: $TASK"
|
| 66 |
+
echo "Groups: $NUM_GROUPS"
|
| 67 |
+
echo "Horizon: 16 (vs baseline 4)"
|
| 68 |
+
echo "Demo: $DEMO_PATH"
|
| 69 |
+
echo "Output: $OUT_DIR"
|
| 70 |
+
echo "=================================================="
|
| 71 |
+
|
| 72 |
+
apptainer exec --nv --env "$ENVS" \
|
| 73 |
+
"$SIF" "$PYTHON" scripts/generate_maniskill_lattice.py \
|
| 74 |
+
--demo "$DEMO_PATH" \
|
| 75 |
+
--out "$OUT_DIR" \
|
| 76 |
+
--env-id "$TASK" \
|
| 77 |
+
--num-groups "$NUM_GROUPS" \
|
| 78 |
+
--k 16 \
|
| 79 |
+
--horizon 16 \
|
| 80 |
+
--seed 0 \
|
| 81 |
+
--shard-size 1024 \
|
| 82 |
+
--sim-backend physx_cuda:0 \
|
| 83 |
+
--render-backend cpu \
|
| 84 |
+
--state-storage archive
|
| 85 |
+
|
| 86 |
+
echo ""
|
| 87 |
+
echo "✅ $TASK generation complete"
|
scripts/slurm/generate_cil_array.sbatch
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=${DOVLA_JOB_NAME:-dovla_cil_gen}
|
| 3 |
+
#SBATCH --partition=${DOVLA_PARTITION:-compute}
|
| 4 |
+
#SBATCH --array=${DOVLA_ARRAY:-0-9}
|
| 5 |
+
#SBATCH --nodes=1
|
| 6 |
+
#SBATCH --ntasks=1
|
| 7 |
+
#SBATCH --cpus-per-task=${DOVLA_CPUS_PER_TASK:-8}
|
| 8 |
+
#SBATCH --gres=gpu:${DOVLA_GPUS_PER_TASK:-0}
|
| 9 |
+
#SBATCH --mem=${DOVLA_MEM:-32G}
|
| 10 |
+
#SBATCH --time=${DOVLA_TIME:-12:00:00}
|
| 11 |
+
#SBATCH --output=${DOVLA_LOG_DIR:-logs/slurm}/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=${DOVLA_LOG_DIR:-logs/slurm}/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 17 |
+
VENV_PATH="${VENV_PATH:-$PROJECT_DIR/.venv}"
|
| 18 |
+
TASKS_PATH="${TASKS_PATH:-$PROJECT_DIR/data/tasks.jsonl}"
|
| 19 |
+
OUT_ROOT="${OUT_ROOT:-$PROJECT_DIR/data/cil_array}"
|
| 20 |
+
BACKEND="${BACKEND:-toy}"
|
| 21 |
+
NUM_WORKERS="${NUM_WORKERS:-4}"
|
| 22 |
+
STATES_PER_TASK="${STATES_PER_TASK:-1000}"
|
| 23 |
+
K="${K:-32}"
|
| 24 |
+
SHARD_SIZE="${SHARD_SIZE:-10000}"
|
| 25 |
+
SEED_BASE="${SEED_BASE:-0}"
|
| 26 |
+
RAY_ADDRESS="${RAY_ADDRESS:-}"
|
| 27 |
+
RESUME_FLAG="${RESUME_FLAG:-}"
|
| 28 |
+
|
| 29 |
+
mkdir -p "${DOVLA_LOG_DIR:-logs/slurm}" "$OUT_ROOT"
|
| 30 |
+
cd "$PROJECT_DIR"
|
| 31 |
+
|
| 32 |
+
if [ -f "$VENV_PATH/bin/activate" ]; then
|
| 33 |
+
# shellcheck disable=SC1091
|
| 34 |
+
source "$VENV_PATH/bin/activate"
|
| 35 |
+
fi
|
| 36 |
+
|
| 37 |
+
export OPENCLAUDE_BASE_URL="${OPENCLAUDE_BASE_URL:-https://open-claude.com/v1}"
|
| 38 |
+
export OPENCLAUDE_MODEL="${OPENCLAUDE_MODEL:-<model>}"
|
| 39 |
+
# Set OPENCLAUDE_API_KEY in the job environment or scheduler secret store. Do not echo it.
|
| 40 |
+
|
| 41 |
+
SEED=$((SEED_BASE + SLURM_ARRAY_TASK_ID))
|
| 42 |
+
OUT_DIR="$OUT_ROOT/part_${SLURM_ARRAY_TASK_ID}"
|
| 43 |
+
|
| 44 |
+
CMD=(
|
| 45 |
+
python scripts/generate_cil_distributed.py
|
| 46 |
+
--backend "$BACKEND"
|
| 47 |
+
--tasks "$TASKS_PATH"
|
| 48 |
+
--out "$OUT_DIR"
|
| 49 |
+
--num-workers "$NUM_WORKERS"
|
| 50 |
+
--num-states-per-task "$STATES_PER_TASK"
|
| 51 |
+
--k "$K"
|
| 52 |
+
--seed "$SEED"
|
| 53 |
+
--shard-size "$SHARD_SIZE"
|
| 54 |
+
)
|
| 55 |
+
|
| 56 |
+
if [ -n "$RAY_ADDRESS" ]; then
|
| 57 |
+
CMD+=(--ray-address "$RAY_ADDRESS")
|
| 58 |
+
fi
|
| 59 |
+
if [ -n "$RESUME_FLAG" ]; then
|
| 60 |
+
CMD+=(--resume)
|
| 61 |
+
fi
|
| 62 |
+
|
| 63 |
+
"${CMD[@]}"
|
scripts/slurm/generate_embeddings.sbatch
ADDED
|
@@ -0,0 +1,28 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=gen_embeddings
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --mem=16000M
|
| 7 |
+
#SBATCH --time=1:00:00
|
| 8 |
+
#SBATCH --output=logs/gen_embeddings_%A.out
|
| 9 |
+
#SBATCH --error=logs/gen_embeddings_%A.err
|
| 10 |
+
|
| 11 |
+
set -euo pipefail
|
| 12 |
+
|
| 13 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 14 |
+
cd "$PROJECT_DIR"
|
| 15 |
+
|
| 16 |
+
source .venv/bin/activate
|
| 17 |
+
|
| 18 |
+
echo "=== Generating Instruction Embeddings (Fast Parallel) ==="
|
| 19 |
+
echo "Using 8 CPU cores for parallel encoding"
|
| 20 |
+
echo ""
|
| 21 |
+
|
| 22 |
+
python scripts/generate_instruction_embeddings.py \
|
| 23 |
+
--dataset /scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection \
|
| 24 |
+
--output /scratch/$USER/dovla/experiments/instruction_embeddings.pkl \
|
| 25 |
+
--cache-dir /scratch/$USER/dovla/experiments/embedding_cache
|
| 26 |
+
|
| 27 |
+
echo ""
|
| 28 |
+
echo "✅ Embeddings generated successfully"
|
scripts/slurm/horizon_sweep_pickcube.sbatch
ADDED
|
@@ -0,0 +1,88 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_horizon_sweep
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=01:30:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# DECISIVE EXPERIMENT: Does action horizon raise the oracle ceiling?
|
| 16 |
+
# Generates PickCube CIL at horizon {4, 8, 16, 32}, measures oracle ceiling each.
|
| 17 |
+
# Baseline (horizon=4) oracle for PickCube = 37.4%.
|
| 18 |
+
|
| 19 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 20 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 21 |
+
SIF="$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif"
|
| 22 |
+
PYTHON="$SCRATCH_ROOT/envs/maniskill/bin/python"
|
| 23 |
+
NATIVE_LIBS="$SCRATCH_ROOT/native_libs/lib"
|
| 24 |
+
CPU_RENDER_LIBS="$SCRATCH_ROOT/cpu_render_libs"
|
| 25 |
+
CA_BUNDLE="$SCRATCH_ROOT/ca-bundle.crt"
|
| 26 |
+
VULKAN_ICD="$CPU_RENDER_LIBS/share/vulkan/icd.d/lvp_icd.x86_64.json"
|
| 27 |
+
DEMO="$SCRATCH_ROOT/maniskill_data/demos/PickCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 28 |
+
OUT_ROOT="${OUT_ROOT:-$SCRATCH_ROOT/experiments/horizon_sweep_pickcube}"
|
| 29 |
+
RUNTIME_DIR="/tmp/$USER/dovla-runtime-$SLURM_JOB_ID"
|
| 30 |
+
CACHE_DIR="/tmp/$USER/dovla-mesa-$SLURM_JOB_ID"
|
| 31 |
+
|
| 32 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 33 |
+
cd "$PROJECT_DIR"
|
| 34 |
+
mkdir -p outputs/hpc/logs "$OUT_ROOT" "$RUNTIME_DIR" "$CACHE_DIR"
|
| 35 |
+
chmod 700 "$RUNTIME_DIR"
|
| 36 |
+
|
| 37 |
+
export OMP_NUM_THREADS=1 OPENBLAS_NUM_THREADS=1 MKL_NUM_THREADS=1 LP_NUM_THREADS=1
|
| 38 |
+
|
| 39 |
+
ENVS="LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS:/.singularity.d/libs,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=1,SSL_CERT_FILE=$CA_BUNDLE,REQUESTS_CA_BUNDLE=$CA_BUNDLE,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1"
|
| 40 |
+
|
| 41 |
+
for H in 4 8 16 32; do
|
| 42 |
+
OUT_DIR="$OUT_ROOT/h${H}"
|
| 43 |
+
echo "=================================================="
|
| 44 |
+
echo "Generating PickCube horizon=$H, 200 groups, K=16"
|
| 45 |
+
echo "=================================================="
|
| 46 |
+
apptainer exec --nv --env "$ENVS" \
|
| 47 |
+
"$SIF" "$PYTHON" scripts/generate_maniskill_lattice.py \
|
| 48 |
+
--demo "$DEMO" \
|
| 49 |
+
--out "$OUT_DIR" \
|
| 50 |
+
--env-id PickCube-v1 \
|
| 51 |
+
--num-groups 200 \
|
| 52 |
+
--k 16 \
|
| 53 |
+
--horizon "$H" \
|
| 54 |
+
--seed 0 \
|
| 55 |
+
--shard-size 1024 \
|
| 56 |
+
--sim-backend physx_cuda:0 \
|
| 57 |
+
--render-backend cpu \
|
| 58 |
+
--state-storage archive
|
| 59 |
+
done
|
| 60 |
+
|
| 61 |
+
echo ""
|
| 62 |
+
echo "=================================================="
|
| 63 |
+
echo "ORACLE CEILING BY HORIZON"
|
| 64 |
+
echo "=================================================="
|
| 65 |
+
apptainer exec --nv --env "$ENVS" "$SIF" "$PYTHON" - <<'PY'
|
| 66 |
+
import sys; sys.path.insert(0,'.')
|
| 67 |
+
from dovla_cil.data.datasets import CILDataset
|
| 68 |
+
import os
|
| 69 |
+
root=os.path.expandvars("/scratch/$USER/dovla/experiments/horizon_sweep_pickcube")
|
| 70 |
+
print(f"{'horizon':>8} {'groups':>7} {'oracle':>8} {'expert':>8} {'mean_reward_spread':>18}")
|
| 71 |
+
for H in [4,8,16,32]:
|
| 72 |
+
d=os.path.join(root,f"h{H}")
|
| 73 |
+
try:
|
| 74 |
+
ds=CILDataset(d)
|
| 75 |
+
except Exception as e:
|
| 76 |
+
print(f"{H:>8} ERROR: {e}"); continue
|
| 77 |
+
n=len(ds.group_ids); orac=0; exp=0; spreads=[]
|
| 78 |
+
for gid in ds.group_ids:
|
| 79 |
+
recs=ds.get_group(gid)
|
| 80 |
+
if any(r.reward.terminal_success for r in recs): orac+=1
|
| 81 |
+
if any(r.candidate_type=='expert' and r.reward.terminal_success for r in recs): exp+=1
|
| 82 |
+
scores=[r.reward.score for r in recs]
|
| 83 |
+
spreads.append(max(scores)-min(scores))
|
| 84 |
+
ms=sum(spreads)/len(spreads) if spreads else 0
|
| 85 |
+
print(f"{H:>8} {n:>7} {orac/n:>8.4f} {exp/n:>8.4f} {ms:>18.4f}")
|
| 86 |
+
print()
|
| 87 |
+
print("Baseline reference: horizon=4 PickCube oracle in full collection = 0.3740")
|
| 88 |
+
PY
|
scripts/slurm/install_smolvla_env.sbatch
ADDED
|
@@ -0,0 +1,104 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_smolvla_env
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=2
|
| 7 |
+
#SBATCH --mem=8G
|
| 8 |
+
#SBATCH --time=00:30:00
|
| 9 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 10 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 11 |
+
|
| 12 |
+
set -euo pipefail
|
| 13 |
+
|
| 14 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 15 |
+
SCRATCH_ROOT="${SCRATCH_ROOT:-/scratch/$USER/dovla}"
|
| 16 |
+
CONTAINER="${CONTAINER:-$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif}"
|
| 17 |
+
ENV_DIR="${ENV_DIR:-$SCRATCH_ROOT/envs/smolvla}"
|
| 18 |
+
LEROBOT_WHEEL="${LEROBOT_WHEEL:-$SCRATCH_ROOT/wheels/lerobot-0.4.3-py3-none-any.whl}"
|
| 19 |
+
DRACCUS_WHEEL="${DRACCUS_WHEEL:-$SCRATCH_ROOT/wheels/draccus-0.10.0-py3-none-any.whl}"
|
| 20 |
+
PYYAML_INCLUDE_WHEEL="${PYYAML_INCLUDE_WHEEL:-$SCRATCH_ROOT/wheels/pyyaml_include-1.4.1-py3-none-any.whl}"
|
| 21 |
+
PYARROW_WHEEL="${PYARROW_WHEEL:-$SCRATCH_ROOT/wheels/pyarrow-17.0.0-cp311-cp311-linux_x86_64.whl}"
|
| 22 |
+
DATASETS_WHEEL="${DATASETS_WHEEL:-/cvmfs/soft.computecanada.ca/custom/python/wheelhouse/generic/datasets-4.0.0+computecanada-py3-none-any.whl}"
|
| 23 |
+
WHEELHOUSE_ARCH="${WHEELHOUSE_ARCH:-/cvmfs/soft.computecanada.ca/custom/python/wheelhouse/gentoo2023/x86-64-v3}"
|
| 24 |
+
WHEELHOUSE_GENERIC="${WHEELHOUSE_GENERIC:-/cvmfs/soft.computecanada.ca/custom/python/wheelhouse/gentoo2023/generic}"
|
| 25 |
+
|
| 26 |
+
cd "$PROJECT_DIR"
|
| 27 |
+
mkdir -p outputs/hpc/logs "$SCRATCH_ROOT/envs"
|
| 28 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 29 |
+
|
| 30 |
+
for WHEEL in \
|
| 31 |
+
"$LEROBOT_WHEEL" \
|
| 32 |
+
"$DRACCUS_WHEEL" \
|
| 33 |
+
"$PYYAML_INCLUDE_WHEEL" \
|
| 34 |
+
"$PYARROW_WHEEL" \
|
| 35 |
+
"$DATASETS_WHEEL"; do
|
| 36 |
+
if [[ ! -f "$WHEEL" ]]; then
|
| 37 |
+
echo "Missing pinned runtime wheel: $WHEEL" >&2
|
| 38 |
+
echo "Stage all pinned wheels before submitting this offline job." >&2
|
| 39 |
+
exit 2
|
| 40 |
+
fi
|
| 41 |
+
done
|
| 42 |
+
|
| 43 |
+
if [[ ! -x "$ENV_DIR/bin/python" ]]; then
|
| 44 |
+
apptainer exec \
|
| 45 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 46 |
+
"$CONTAINER" \
|
| 47 |
+
/opt/conda/bin/python -m venv --system-site-packages "$ENV_DIR"
|
| 48 |
+
fi
|
| 49 |
+
|
| 50 |
+
apptainer exec \
|
| 51 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 52 |
+
-B "$PROJECT_DIR:$PROJECT_DIR" \
|
| 53 |
+
-B /cvmfs:/cvmfs \
|
| 54 |
+
"$CONTAINER" \
|
| 55 |
+
"$ENV_DIR/bin/python" -c \
|
| 56 |
+
"from itertools import islice; from packaging.tags import sys_tags; print('supported_tags', [str(tag) for tag in islice(sys_tags(), 12)])"
|
| 57 |
+
|
| 58 |
+
apptainer exec \
|
| 59 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 60 |
+
-B "$PROJECT_DIR:$PROJECT_DIR" \
|
| 61 |
+
-B /cvmfs:/cvmfs \
|
| 62 |
+
"$CONTAINER" \
|
| 63 |
+
"$ENV_DIR/bin/python" -m pip install \
|
| 64 |
+
--no-index \
|
| 65 |
+
--find-links "$WHEELHOUSE_ARCH" \
|
| 66 |
+
--find-links "$WHEELHOUSE_GENERIC" \
|
| 67 |
+
"transformers==4.57.6+computecanada" \
|
| 68 |
+
"huggingface-hub==0.35.3+computecanada" \
|
| 69 |
+
"accelerate==1.10.1+computecanada" \
|
| 70 |
+
"num2words==0.5.14+computecanada" \
|
| 71 |
+
"typing-inspect==0.9.0+computecanada" \
|
| 72 |
+
"mergedeep==1.3.4+computecanada" \
|
| 73 |
+
"toml==0.10.2+computecanada" \
|
| 74 |
+
"einops==0.8.1+computecanada" \
|
| 75 |
+
"dill==0.3.8+computecanada" \
|
| 76 |
+
"multiprocess==0.70.16+computecanada" \
|
| 77 |
+
"xxhash==3.5.0+computecanada" \
|
| 78 |
+
"pandas==2.2.3+computecanada" \
|
| 79 |
+
"fsspec==2025.3.0+computecanada" \
|
| 80 |
+
"setuptools==80.9.0+computecanada" \
|
| 81 |
+
"imageio==2.37.0+computecanada" \
|
| 82 |
+
"imageio-ffmpeg==0.6.0+computecanada"
|
| 83 |
+
|
| 84 |
+
apptainer exec \
|
| 85 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 86 |
+
-B "$PROJECT_DIR:$PROJECT_DIR" \
|
| 87 |
+
-B /cvmfs:/cvmfs \
|
| 88 |
+
"$CONTAINER" \
|
| 89 |
+
"$ENV_DIR/bin/python" -m pip install \
|
| 90 |
+
--no-index \
|
| 91 |
+
--no-deps \
|
| 92 |
+
"$PYYAML_INCLUDE_WHEEL" \
|
| 93 |
+
"$PYARROW_WHEEL" \
|
| 94 |
+
"$DATASETS_WHEEL" \
|
| 95 |
+
"$DRACCUS_WHEEL" \
|
| 96 |
+
"$LEROBOT_WHEEL"
|
| 97 |
+
|
| 98 |
+
apptainer exec \
|
| 99 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 100 |
+
-B "$PROJECT_DIR:$PROJECT_DIR" \
|
| 101 |
+
--env "PYTHONPATH=$PROJECT_DIR" \
|
| 102 |
+
"$CONTAINER" \
|
| 103 |
+
"$ENV_DIR/bin/python" -c \
|
| 104 |
+
"import accelerate, datasets, importlib.util, lerobot, pyarrow, transformers; from dovla_cil.eval.smolvla_runtime import import_smolvla_classes; SmolVLAPolicy, SmolVLAConfig = import_smolvla_classes(); print('lerobot', lerobot.__version__); print('transformers', transformers.__version__); print('accelerate', accelerate.__version__); print('datasets', datasets.__version__); print('pyarrow', pyarrow.__version__); print('policy_import', SmolVLAPolicy.__name__, SmolVLAConfig.__name__); [print(name, bool(importlib.util.find_spec(name))) for name in ('draccus', 'typing_inspect', 'gymnasium', 'einops', 'safetensors')]"
|
scripts/slurm/make_maniskill_collection.sbatch
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_collect
|
| 3 |
+
#SBATCH --account=def-yalda_cpu
|
| 4 |
+
#SBATCH --partition=cpubase_bycore_b1
|
| 5 |
+
#SBATCH --nodes=1
|
| 6 |
+
#SBATCH --ntasks=1
|
| 7 |
+
#SBATCH --cpus-per-task=2
|
| 8 |
+
#SBATCH --mem=8G
|
| 9 |
+
#SBATCH --time=00:20:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
PICKCUBE_DATA="${PICKCUBE_DATA:?Set PICKCUBE_DATA}"
|
| 17 |
+
MULTITASK_ROOT="${MULTITASK_ROOT:?Set MULTITASK_ROOT}"
|
| 18 |
+
COLLECTION_OUT="${COLLECTION_OUT:?Set COLLECTION_OUT}"
|
| 19 |
+
COLLECTION_NAME="${COLLECTION_NAME:-maniskill-six-task-k16}"
|
| 20 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 21 |
+
|
| 22 |
+
cd "$PROJECT_DIR"
|
| 23 |
+
"$PYTHON" scripts/make_cil_collection.py \
|
| 24 |
+
--name "$COLLECTION_NAME" \
|
| 25 |
+
--out "$COLLECTION_OUT" \
|
| 26 |
+
--sources \
|
| 27 |
+
"$PICKCUBE_DATA" \
|
| 28 |
+
"$MULTITASK_ROOT/PushCube-v1" \
|
| 29 |
+
"$MULTITASK_ROOT/PullCube-v1" \
|
| 30 |
+
"$MULTITASK_ROOT/StackCube-v1" \
|
| 31 |
+
"$MULTITASK_ROOT/LiftPegUpright-v1" \
|
| 32 |
+
"$MULTITASK_ROOT/PegInsertionSide-v1"
|
scripts/slurm/maniskill_lattice_debug.sbatch
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_debug
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
# Physics uses CUDA; state-mode material creation uses the CPU Vulkan renderer to avoid
|
| 8 |
+
# Vulkan/CUDA device-ordinal mismatches on shared four-GPU nodes.
|
| 9 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 10 |
+
#SBATCH --mem=24G
|
| 11 |
+
#SBATCH --time=00:20:00
|
| 12 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 13 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 14 |
+
|
| 15 |
+
set -euo pipefail
|
| 16 |
+
|
| 17 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 18 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 19 |
+
SIF="$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif"
|
| 20 |
+
PYTHON="$SCRATCH_ROOT/envs/maniskill/bin/python"
|
| 21 |
+
NATIVE_LIBS="$SCRATCH_ROOT/native_libs/lib"
|
| 22 |
+
CPU_RENDER_LIBS="$SCRATCH_ROOT/cpu_render_libs"
|
| 23 |
+
CA_BUNDLE="$SCRATCH_ROOT/ca-bundle.crt"
|
| 24 |
+
VULKAN_ICD="$CPU_RENDER_LIBS/share/vulkan/icd.d/lvp_icd.x86_64.json"
|
| 25 |
+
DEMO="$SCRATCH_ROOT/maniskill_data/demos/PickCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 26 |
+
OUT_DIR="${OUT_DIR:-$PROJECT_DIR/outputs/hpc/maniskill_debug_cil}"
|
| 27 |
+
RUNTIME_DIR="/tmp/$USER/dovla-runtime-$SLURM_JOB_ID"
|
| 28 |
+
CACHE_DIR="/tmp/$USER/dovla-mesa-$SLURM_JOB_ID"
|
| 29 |
+
|
| 30 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 31 |
+
cd "$PROJECT_DIR"
|
| 32 |
+
mkdir -p outputs/hpc/logs "$OUT_DIR" "$RUNTIME_DIR" "$CACHE_DIR"
|
| 33 |
+
chmod 700 "$RUNTIME_DIR"
|
| 34 |
+
|
| 35 |
+
export OMP_NUM_THREADS=1
|
| 36 |
+
export OPENBLAS_NUM_THREADS=1
|
| 37 |
+
export MKL_NUM_THREADS=1
|
| 38 |
+
export LP_NUM_THREADS=1
|
| 39 |
+
|
| 40 |
+
apptainer exec --nv \
|
| 41 |
+
--env "LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS:/.singularity.d/libs,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=1,SSL_CERT_FILE=$CA_BUNDLE,REQUESTS_CA_BUNDLE=$CA_BUNDLE,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1" \
|
| 42 |
+
"$SIF" "$PYTHON" - <<'PY'
|
| 43 |
+
import gymnasium as gym
|
| 44 |
+
import mani_skill
|
| 45 |
+
import os
|
| 46 |
+
import torch
|
| 47 |
+
|
| 48 |
+
print("torch", torch.__version__, "cuda", torch.cuda.is_available(), torch.cuda.get_device_name(0))
|
| 49 |
+
print("vulkan_icd", os.environ.get("VK_ICD_FILENAMES"), "cuda_visible", os.environ.get("CUDA_VISIBLE_DEVICES"))
|
| 50 |
+
env = gym.make(
|
| 51 |
+
"PickCube-v1",
|
| 52 |
+
num_envs=1,
|
| 53 |
+
obs_mode="state",
|
| 54 |
+
control_mode="pd_ee_delta_pose",
|
| 55 |
+
render_mode=None,
|
| 56 |
+
sim_backend="physx_cuda:0",
|
| 57 |
+
render_backend="cpu",
|
| 58 |
+
)
|
| 59 |
+
env.reset(seed=7)
|
| 60 |
+
state = {
|
| 61 |
+
section: {name: value.clone() for name, value in values.items()}
|
| 62 |
+
for section, values in env.unwrapped.get_state_dict().items()
|
| 63 |
+
}
|
| 64 |
+
action = torch.zeros((1, 7), dtype=torch.float32, device=env.unwrapped.device)
|
| 65 |
+
env.unwrapped.set_state_dict(state)
|
| 66 |
+
env.unwrapped.agent.controller.reset()
|
| 67 |
+
restored = env.unwrapped.get_state_dict()
|
| 68 |
+
max_error = max(
|
| 69 |
+
float(torch.max(torch.abs(state[section][name] - restored[section][name])).cpu())
|
| 70 |
+
for section in state
|
| 71 |
+
for name in state[section]
|
| 72 |
+
)
|
| 73 |
+
print("state_restore_max_error", max_error)
|
| 74 |
+
assert max_error <= 1e-6
|
| 75 |
+
|
| 76 |
+
env.unwrapped.step(action)
|
| 77 |
+
next_state_1 = {
|
| 78 |
+
section: {name: value.clone() for name, value in values.items()}
|
| 79 |
+
for section, values in env.unwrapped.get_state_dict().items()
|
| 80 |
+
}
|
| 81 |
+
env.unwrapped.set_state_dict(state)
|
| 82 |
+
env.unwrapped.agent.controller.reset()
|
| 83 |
+
env.unwrapped.step(action)
|
| 84 |
+
next_state_2 = env.unwrapped.get_state_dict()
|
| 85 |
+
branch_error = max(
|
| 86 |
+
float(torch.max(torch.abs(next_state_1[section][name] - next_state_2[section][name])).cpu())
|
| 87 |
+
for section in next_state_1
|
| 88 |
+
for name in next_state_1[section]
|
| 89 |
+
)
|
| 90 |
+
print("deterministic_branch_max_error", branch_error)
|
| 91 |
+
assert branch_error <= 1e-5
|
| 92 |
+
env.close()
|
| 93 |
+
PY
|
| 94 |
+
|
| 95 |
+
apptainer exec --nv \
|
| 96 |
+
--env "LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS:/.singularity.d/libs,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=1,SSL_CERT_FILE=$CA_BUNDLE,REQUESTS_CA_BUNDLE=$CA_BUNDLE,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1" \
|
| 97 |
+
"$SIF" "$PYTHON" scripts/generate_maniskill_lattice.py \
|
| 98 |
+
--demo "$DEMO" \
|
| 99 |
+
--out "$OUT_DIR" \
|
| 100 |
+
--num-groups 8 \
|
| 101 |
+
--k 4 \
|
| 102 |
+
--horizon 4 \
|
| 103 |
+
--seed 0 \
|
| 104 |
+
--shard-size 32 \
|
| 105 |
+
--sim-backend physx_cuda:0 \
|
| 106 |
+
--render-backend cpu \
|
| 107 |
+
--state-storage archive
|
| 108 |
+
|
| 109 |
+
apptainer exec --nv \
|
| 110 |
+
--env "LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS:/.singularity.d/libs,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=1,SSL_CERT_FILE=$CA_BUNDLE,REQUESTS_CA_BUNDLE=$CA_BUNDLE,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1" \
|
| 111 |
+
"$SIF" "$PYTHON" scripts/inspect_shard.py "$OUT_DIR/manifest.json"
|
scripts/slurm/maniskill_lattice_full.sbatch
ADDED
|
@@ -0,0 +1,80 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_full
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=32G
|
| 9 |
+
#SBATCH --time=02:00:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 17 |
+
SIF="$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif"
|
| 18 |
+
PYTHON="$SCRATCH_ROOT/envs/maniskill/bin/python"
|
| 19 |
+
NATIVE_LIBS="$SCRATCH_ROOT/native_libs/lib"
|
| 20 |
+
CPU_RENDER_LIBS="$SCRATCH_ROOT/cpu_render_libs"
|
| 21 |
+
CA_BUNDLE="$SCRATCH_ROOT/ca-bundle.crt"
|
| 22 |
+
VULKAN_ICD="$CPU_RENDER_LIBS/share/vulkan/icd.d/lvp_icd.x86_64.json"
|
| 23 |
+
DEMO="${DEMO:-$SCRATCH_ROOT/maniskill_data/demos/PickCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5}"
|
| 24 |
+
ENV_ID="${ENV_ID:-PickCube-v1}"
|
| 25 |
+
CONTROL_MODE="${CONTROL_MODE:-pd_ee_delta_pose}"
|
| 26 |
+
|
| 27 |
+
NUM_GROUPS="${NUM_GROUPS:-1000}"
|
| 28 |
+
GROUP_OFFSET="${GROUP_OFFSET:-0}"
|
| 29 |
+
K="${K:-16}"
|
| 30 |
+
HORIZON="${HORIZON:-4}"
|
| 31 |
+
SEED="${SEED:-0}"
|
| 32 |
+
SHARD_SIZE="${SHARD_SIZE:-2048}"
|
| 33 |
+
STATE_BATCH_SIZE="${STATE_BATCH_SIZE:-16}"
|
| 34 |
+
OBS_MODE="${OBS_MODE:-state}"
|
| 35 |
+
IMAGE_QUALITY="${IMAGE_QUALITY:-90}"
|
| 36 |
+
CANDIDATE_MODE="${CANDIDATE_MODE:-structured}"
|
| 37 |
+
OUT_DIR="${OUT_DIR:-$PROJECT_DIR/outputs/hpc/maniskill_full_k${K}_n${NUM_GROUPS}_seed${SEED}}"
|
| 38 |
+
RUNTIME_DIR="/tmp/$USER/dovla-runtime-$SLURM_JOB_ID"
|
| 39 |
+
CACHE_DIR="/tmp/$USER/dovla-mesa-$SLURM_JOB_ID"
|
| 40 |
+
|
| 41 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 42 |
+
cd "$PROJECT_DIR"
|
| 43 |
+
mkdir -p outputs/hpc/logs "$OUT_DIR" "$RUNTIME_DIR" "$CACHE_DIR"
|
| 44 |
+
chmod 700 "$RUNTIME_DIR"
|
| 45 |
+
|
| 46 |
+
export OMP_NUM_THREADS=1
|
| 47 |
+
export OPENBLAS_NUM_THREADS=1
|
| 48 |
+
export MKL_NUM_THREADS=1
|
| 49 |
+
export LP_NUM_THREADS=1
|
| 50 |
+
|
| 51 |
+
if [[ -f "$OUT_DIR/manifest.json" ]]; then
|
| 52 |
+
echo "completed manifest already exists: $OUT_DIR/manifest.json"
|
| 53 |
+
exit 0
|
| 54 |
+
fi
|
| 55 |
+
|
| 56 |
+
apptainer exec --nv \
|
| 57 |
+
--env "LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS:/.singularity.d/libs,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=1,SSL_CERT_FILE=$CA_BUNDLE,REQUESTS_CA_BUNDLE=$CA_BUNDLE,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1" \
|
| 58 |
+
"$SIF" "$PYTHON" scripts/generate_maniskill_lattice.py \
|
| 59 |
+
--demo "$DEMO" \
|
| 60 |
+
--env-id "$ENV_ID" \
|
| 61 |
+
--control-mode "$CONTROL_MODE" \
|
| 62 |
+
--out "$OUT_DIR" \
|
| 63 |
+
--num-groups "$NUM_GROUPS" \
|
| 64 |
+
--group-offset "$GROUP_OFFSET" \
|
| 65 |
+
--k "$K" \
|
| 66 |
+
--horizon "$HORIZON" \
|
| 67 |
+
--seed "$SEED" \
|
| 68 |
+
--shard-size "$SHARD_SIZE" \
|
| 69 |
+
--obs-mode "$OBS_MODE" \
|
| 70 |
+
--image-quality "$IMAGE_QUALITY" \
|
| 71 |
+
--sim-backend physx_cuda:0 \
|
| 72 |
+
--render-backend cpu \
|
| 73 |
+
--parallel-branches \
|
| 74 |
+
--state-batch-size "$STATE_BATCH_SIZE" \
|
| 75 |
+
--state-storage archive \
|
| 76 |
+
--candidate-mode "$CANDIDATE_MODE"
|
| 77 |
+
|
| 78 |
+
apptainer exec \
|
| 79 |
+
--env "OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1" \
|
| 80 |
+
"$SIF" "$PYTHON" scripts/inspect_shard.py "$OUT_DIR/manifest.json"
|
scripts/slurm/maniskill_multitask_pilot.sbatch
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_multi
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=32G
|
| 9 |
+
#SBATCH --time=00:20:00
|
| 10 |
+
#SBATCH --array=0-4%2
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 18 |
+
DEMO_ROOT="$SCRATCH_ROOT/maniskill_multitask_demos"
|
| 19 |
+
MULTITASK_OUT_ROOT="${MULTITASK_OUT_ROOT:-$SCRATCH_ROOT/experiments/maniskill_multitask_pilot}"
|
| 20 |
+
|
| 21 |
+
case "${SLURM_ARRAY_TASK_ID:-0}" in
|
| 22 |
+
0)
|
| 23 |
+
ENV_ID="PushCube-v1"
|
| 24 |
+
CONTROL_MODE="pd_ee_delta_pose"
|
| 25 |
+
DEMO="$DEMO_ROOT/PushCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 26 |
+
;;
|
| 27 |
+
1)
|
| 28 |
+
ENV_ID="PullCube-v1"
|
| 29 |
+
CONTROL_MODE="pd_ee_delta_pose"
|
| 30 |
+
DEMO="$DEMO_ROOT/PullCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 31 |
+
;;
|
| 32 |
+
2)
|
| 33 |
+
ENV_ID="StackCube-v1"
|
| 34 |
+
CONTROL_MODE="pd_ee_delta_pose"
|
| 35 |
+
DEMO="$DEMO_ROOT/StackCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 36 |
+
;;
|
| 37 |
+
3)
|
| 38 |
+
ENV_ID="LiftPegUpright-v1"
|
| 39 |
+
CONTROL_MODE="pd_ee_delta_pose"
|
| 40 |
+
DEMO="$DEMO_ROOT/LiftPegUpright-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 41 |
+
;;
|
| 42 |
+
4)
|
| 43 |
+
ENV_ID="PegInsertionSide-v1"
|
| 44 |
+
CONTROL_MODE="pd_joint_pos"
|
| 45 |
+
DEMO="$DEMO_ROOT/PegInsertionSide-v1/motionplanning/trajectory.h5"
|
| 46 |
+
;;
|
| 47 |
+
*)
|
| 48 |
+
echo "unsupported array index" >&2
|
| 49 |
+
exit 2
|
| 50 |
+
;;
|
| 51 |
+
esac
|
| 52 |
+
|
| 53 |
+
export PROJECT_DIR DEMO ENV_ID CONTROL_MODE
|
| 54 |
+
export NUM_GROUPS="${NUM_GROUPS:-16}"
|
| 55 |
+
export K="${K:-8}"
|
| 56 |
+
export HORIZON="${HORIZON:-4}"
|
| 57 |
+
export STATE_BATCH_SIZE="${STATE_BATCH_SIZE:-8}"
|
| 58 |
+
export OBS_MODE=state
|
| 59 |
+
export SHARD_SIZE="${SHARD_SIZE:-256}"
|
| 60 |
+
export STATE_STORAGE=archive
|
| 61 |
+
export OUT_DIR="$MULTITASK_OUT_ROOT/$ENV_ID"
|
| 62 |
+
|
| 63 |
+
exec bash "$PROJECT_DIR/scripts/slurm/maniskill_lattice_full.sbatch"
|
scripts/slurm/phase_a1_generate_10k.sbatch
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_10k_gen
|
| 3 |
+
#SBATCH --partition=${DOVLA_PARTITION:-compute}
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=16
|
| 7 |
+
#SBATCH --gres=gpu:1
|
| 8 |
+
#SBATCH --mem=64G
|
| 9 |
+
#SBATCH --time=48:00:00
|
| 10 |
+
#SBATCH --output=logs/phase_a_10k_gen_%j.out
|
| 11 |
+
#SBATCH --error=logs/phase_a_10k_gen_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase A1: Generate 10K groups dataset for performance improvement
|
| 16 |
+
# This scales from current 3,500 to 10,000 groups
|
| 17 |
+
# Expected: +5-10% success improvement
|
| 18 |
+
|
| 19 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 20 |
+
cd "$PROJECT_DIR"
|
| 21 |
+
|
| 22 |
+
# Activate environment
|
| 23 |
+
if [ -f ".venv/bin/activate" ]; then
|
| 24 |
+
source .venv/bin/activate
|
| 25 |
+
fi
|
| 26 |
+
|
| 27 |
+
# Configuration
|
| 28 |
+
DEMO_DIR="/scratch/$USER/dovla/demonstrations/maniskill"
|
| 29 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_a_10k_collection"
|
| 30 |
+
K=16
|
| 31 |
+
STATE_BATCH_SIZE=16
|
| 32 |
+
|
| 33 |
+
# Task configuration: 6 tasks with more groups each
|
| 34 |
+
declare -A TASK_GROUPS=(
|
| 35 |
+
["PickCube-v1"]=2000 # Increase from 1000
|
| 36 |
+
["PushCube-v1"]=2000 # Increase from 500
|
| 37 |
+
["PullCube-v1"]=1500 # Increase from 500
|
| 38 |
+
["StackCube-v1"]=1500 # Increase from 500
|
| 39 |
+
["LiftPegUpright-v1"]=1500 # Increase from 500
|
| 40 |
+
["PegInsertionSide-v1"]=1500 # Increase from 500
|
| 41 |
+
)
|
| 42 |
+
|
| 43 |
+
mkdir -p "$OUT_DIR" logs
|
| 44 |
+
|
| 45 |
+
echo "=== Phase A1: Generating 10K Group Collection ==="
|
| 46 |
+
echo "Target: 10,000 groups, 160,000 records (K=$K)"
|
| 47 |
+
echo "Expected improvement: +5-10% policy success"
|
| 48 |
+
echo ""
|
| 49 |
+
|
| 50 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 51 |
+
NUM_GROUPS="${TASK_GROUPS[$TASK]}"
|
| 52 |
+
DEMO_FILE="$DEMO_DIR/${TASK}.h5"
|
| 53 |
+
TASK_OUT="$OUT_DIR/${TASK}_k${K}_n${NUM_GROUPS}"
|
| 54 |
+
|
| 55 |
+
if [ ! -f "$DEMO_FILE" ]; then
|
| 56 |
+
echo "⚠️ Demo file not found: $DEMO_FILE"
|
| 57 |
+
echo " Skipping $TASK"
|
| 58 |
+
continue
|
| 59 |
+
fi
|
| 60 |
+
|
| 61 |
+
echo "Generating $TASK: $NUM_GROUPS groups..."
|
| 62 |
+
|
| 63 |
+
python scripts/generate_maniskill_lattice.py \
|
| 64 |
+
--demo "$DEMO_FILE" \
|
| 65 |
+
--env-id "$TASK" \
|
| 66 |
+
--control-mode pd_ee_delta_pose \
|
| 67 |
+
--out "$TASK_OUT" \
|
| 68 |
+
--num-groups "$NUM_GROUPS" \
|
| 69 |
+
--k "$K" \
|
| 70 |
+
--state-batch-size "$STATE_BATCH_SIZE" \
|
| 71 |
+
--seed 42 \
|
| 72 |
+
--pre-success-only
|
| 73 |
+
|
| 74 |
+
if [ $? -eq 0 ]; then
|
| 75 |
+
echo "✅ $TASK complete: $NUM_GROUPS groups"
|
| 76 |
+
else
|
| 77 |
+
echo "❌ $TASK failed"
|
| 78 |
+
exit 1
|
| 79 |
+
fi
|
| 80 |
+
echo ""
|
| 81 |
+
done
|
| 82 |
+
|
| 83 |
+
echo "=== Merging into unified collection ==="
|
| 84 |
+
|
| 85 |
+
python scripts/make_cil_collection.py \
|
| 86 |
+
--source-dirs "$OUT_DIR"/*/ \
|
| 87 |
+
--out "$OUT_DIR/merged_10k" \
|
| 88 |
+
--name "phase_a_10k_collection"
|
| 89 |
+
|
| 90 |
+
echo "✅ Phase A1 complete: 10K group collection ready"
|
| 91 |
+
echo " Location: $OUT_DIR/merged_10k"
|
| 92 |
+
echo ""
|
| 93 |
+
echo "Next: Run phase_a2_train_large_model.sbatch"
|
scripts/slurm/phase_a1_generate_10k_enhanced.sbatch
ADDED
|
@@ -0,0 +1,120 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_10k_gen
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=16
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=96:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_a1_10k_gen_%j.out
|
| 10 |
+
#SBATCH --error=logs/phase_a1_10k_gen_%j.err
|
| 11 |
+
|
| 12 |
+
set -euo pipefail
|
| 13 |
+
|
| 14 |
+
# Phase A1: Enhanced 10K Generation
|
| 15 |
+
# Target: 50%+ policy success with optimizations
|
| 16 |
+
|
| 17 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 18 |
+
cd "$PROJECT_DIR"
|
| 19 |
+
|
| 20 |
+
source .venv/bin/activate
|
| 21 |
+
|
| 22 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_a1_10k_collection"
|
| 23 |
+
K=16
|
| 24 |
+
STATE_BATCH_SIZE=16
|
| 25 |
+
|
| 26 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 27 |
+
echo "Phase A1: Enhanced 10K Generation for 50%+ Target"
|
| 28 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 29 |
+
echo ""
|
| 30 |
+
echo "Strategy:"
|
| 31 |
+
echo " - 10,000 groups (vs 3,500 current)"
|
| 32 |
+
echo " - 160,000 records total"
|
| 33 |
+
echo " - K=16 interventions per group"
|
| 34 |
+
echo " - Optimized for diverse counterfactuals"
|
| 35 |
+
echo ""
|
| 36 |
+
echo "Expected outcome: 42-50% policy success"
|
| 37 |
+
echo ""
|
| 38 |
+
|
| 39 |
+
# Task distribution (balanced across difficulty)
|
| 40 |
+
declare -A TASK_GROUPS=(
|
| 41 |
+
["PickCube-v1"]=1800 # Easy
|
| 42 |
+
["PushCube-v1"]=1800 # Easy
|
| 43 |
+
["PullCube-v1"]=1600 # Medium
|
| 44 |
+
["StackCube-v1"]=1600 # Medium-Hard
|
| 45 |
+
["LiftPegUpright-v1"]=1600 # Medium-Hard
|
| 46 |
+
["PegInsertionSide-v1"]=1600 # Hard
|
| 47 |
+
)
|
| 48 |
+
|
| 49 |
+
TOTAL_GROUPS=0
|
| 50 |
+
for count in "${TASK_GROUPS[@]}"; do
|
| 51 |
+
TOTAL_GROUPS=$((TOTAL_GROUPS + count))
|
| 52 |
+
done
|
| 53 |
+
|
| 54 |
+
echo "Task distribution (total: $TOTAL_GROUPS groups):"
|
| 55 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 56 |
+
echo " ${TASK}: ${TASK_GROUPS[$TASK]} groups"
|
| 57 |
+
done
|
| 58 |
+
echo ""
|
| 59 |
+
|
| 60 |
+
# Generate each task
|
| 61 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 62 |
+
NUM_GROUPS="${TASK_GROUPS[$TASK]}"
|
| 63 |
+
|
| 64 |
+
TASK_OUT="$OUT_DIR/${TASK}_k${K}_n${NUM_GROUPS}"
|
| 65 |
+
|
| 66 |
+
if [ -d "$TASK_OUT/merged" ]; then
|
| 67 |
+
echo "✓ $TASK already generated, skipping"
|
| 68 |
+
continue
|
| 69 |
+
fi
|
| 70 |
+
|
| 71 |
+
echo "Generating $TASK: $NUM_GROUPS groups..."
|
| 72 |
+
echo " Start: $(date)"
|
| 73 |
+
|
| 74 |
+
# Determine demo path
|
| 75 |
+
DEMO_PATH="/scratch/$USER/dovla/demos/maniskill/${TASK%.v1}.h5"
|
| 76 |
+
if [ ! -f "$DEMO_PATH" ]; then
|
| 77 |
+
echo " ⚠️ Demo not found at $DEMO_PATH, trying alternate location..."
|
| 78 |
+
DEMO_PATH="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection/${TASK}/demos/demo.h5"
|
| 79 |
+
fi
|
| 80 |
+
|
| 81 |
+
if [ ! -f "$DEMO_PATH" ]; then
|
| 82 |
+
echo " ❌ Demo not found, skipping $TASK"
|
| 83 |
+
continue
|
| 84 |
+
fi
|
| 85 |
+
|
| 86 |
+
python scripts/generate_maniskill_lattice.py \
|
| 87 |
+
--demo "$DEMO_PATH" \
|
| 88 |
+
--env-id "$TASK" \
|
| 89 |
+
--control-mode pd_ee_delta_pose \
|
| 90 |
+
--out "$TASK_OUT" \
|
| 91 |
+
--num-groups "$NUM_GROUPS" \
|
| 92 |
+
--k "$K" \
|
| 93 |
+
--state-batch-size "$STATE_BATCH_SIZE" \
|
| 94 |
+
--seed 42
|
| 95 |
+
|
| 96 |
+
echo " ✅ Complete: $(date)"
|
| 97 |
+
echo ""
|
| 98 |
+
done
|
| 99 |
+
|
| 100 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 101 |
+
echo "Merging all tasks into unified collection"
|
| 102 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 103 |
+
|
| 104 |
+
python scripts/make_cil_collection.py \
|
| 105 |
+
--source-dirs "$OUT_DIR"/*/merged \
|
| 106 |
+
--out "$OUT_DIR/merged_10k" \
|
| 107 |
+
--name "phase_a1_10k_enhanced"
|
| 108 |
+
|
| 109 |
+
echo ""
|
| 110 |
+
echo "✅ Phase A1 Enhanced Generation Complete!"
|
| 111 |
+
echo ""
|
| 112 |
+
echo "Output: $OUT_DIR/merged_10k"
|
| 113 |
+
echo "Total groups: $TOTAL_GROUPS"
|
| 114 |
+
echo "Total records: $((TOTAL_GROUPS * K))"
|
| 115 |
+
echo ""
|
| 116 |
+
echo "Next: Train enhanced model with:"
|
| 117 |
+
echo " - Hidden dim: 512"
|
| 118 |
+
echo " - Epochs: 150"
|
| 119 |
+
echo " - LR: 0.0003"
|
| 120 |
+
echo " - Enhanced loss weights"
|
scripts/slurm/phase_a1_revised_enhanced.sbatch
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_enhanced_train
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_a1_enhanced_single_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/phase_a1_enhanced_single_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase A1-Revised: Enhanced Training on Existing 3.5K Data
|
| 16 |
+
# Target: 45%+ with better training, no new data needed
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_a1_revised_enhanced"
|
| 25 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 26 |
+
|
| 27 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 28 |
+
|
| 29 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 30 |
+
echo "Phase A1-Revised: Enhanced Training (Existing Data)"
|
| 31 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 32 |
+
echo ""
|
| 33 |
+
echo "Strategy: Better training, not more data"
|
| 34 |
+
echo "Seed: $SEED"
|
| 35 |
+
echo "Dataset: 3,500 groups (existing)"
|
| 36 |
+
echo "Model: h=256 (best from Phase A4)"
|
| 37 |
+
echo "Training: 200 epochs with cosine schedule"
|
| 38 |
+
echo ""
|
| 39 |
+
echo "Target: 45%+ policy success"
|
| 40 |
+
echo ""
|
| 41 |
+
|
| 42 |
+
python scripts/train_dovla.py \
|
| 43 |
+
--dataset "$DATASET" \
|
| 44 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 45 |
+
--objective lattice_field \
|
| 46 |
+
--hidden-dim 256 \
|
| 47 |
+
--action-horizon 4 \
|
| 48 |
+
--epochs 200 \
|
| 49 |
+
--batch-groups 16 \
|
| 50 |
+
--records-per-group 8 \
|
| 51 |
+
--lr 0.0003 \
|
| 52 |
+
--weight-decay 0.01 \
|
| 53 |
+
--device auto \
|
| 54 |
+
--seed $SEED \
|
| 55 |
+
--observation-mode state \
|
| 56 |
+
--loss-weight bc=1.0 \
|
| 57 |
+
--loss-weight field_effect=1.5 \
|
| 58 |
+
--loss-weight field_potential=1.0 \
|
| 59 |
+
--loss-weight field_preference=0.8 \
|
| 60 |
+
--loss-weight field_anchor=0.2
|
| 61 |
+
|
| 62 |
+
echo ""
|
| 63 |
+
echo "✅ Phase A1-Revised enhanced training complete (seed $SEED)"
|
| 64 |
+
echo ""
|
| 65 |
+
echo "Next: Evaluate and check if 45%+ achieved"
|
scripts/slurm/phase_a1b_train_enhanced.sbatch
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_enhanced_train
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=120:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_a1b_enhanced_train_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/phase_a1b_enhanced_train_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase A1b: Enhanced Training for 50%+ Target
|
| 16 |
+
|
| 17 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 18 |
+
cd "$PROJECT_DIR"
|
| 19 |
+
|
| 20 |
+
source .venv/bin/activate
|
| 21 |
+
|
| 22 |
+
DATASET="/scratch/$USER/dovla/experiments/phase_a1_10k_collection/merged_10k"
|
| 23 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_a1b_enhanced_model"
|
| 24 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 25 |
+
|
| 26 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 27 |
+
|
| 28 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 29 |
+
echo "Phase A1b: Enhanced Training for 50%+ Target"
|
| 30 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 31 |
+
echo ""
|
| 32 |
+
echo "Seed: $SEED"
|
| 33 |
+
echo "Dataset: 10,000 groups, 160,000 records"
|
| 34 |
+
echo "Model: hidden_dim=512 (optimal size for 10K data)"
|
| 35 |
+
echo "Training: 150 epochs with warmup + decay"
|
| 36 |
+
echo ""
|
| 37 |
+
echo "Target: 50%+ policy success"
|
| 38 |
+
echo ""
|
| 39 |
+
|
| 40 |
+
python scripts/train_dovla.py \
|
| 41 |
+
--dataset "$DATASET" \
|
| 42 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 43 |
+
--objective lattice_field \
|
| 44 |
+
--hidden-dim 512 \
|
| 45 |
+
--action-horizon 4 \
|
| 46 |
+
--epochs 150 \
|
| 47 |
+
--batch-groups 16 \
|
| 48 |
+
--records-per-group 8 \
|
| 49 |
+
--lr 0.0003 \
|
| 50 |
+
--weight-decay 0.01 \
|
| 51 |
+
--device auto \
|
| 52 |
+
--seed $SEED \
|
| 53 |
+
--observation-mode state \
|
| 54 |
+
--loss-weight bc=1.0 \
|
| 55 |
+
--loss-weight field_effect=1.5 \
|
| 56 |
+
--loss-weight field_potential=1.0 \
|
| 57 |
+
--loss-weight field_preference=0.8 \
|
| 58 |
+
--loss-weight field_anchor=0.2
|
| 59 |
+
|
| 60 |
+
echo ""
|
| 61 |
+
echo "✅ Phase A1b enhanced training complete (seed $SEED)"
|
| 62 |
+
echo ""
|
| 63 |
+
echo "Next: Evaluate and compare with 38.43% baseline"
|
scripts/slurm/phase_a2_train_large_model.sbatch
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_large_train
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=72:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_a2_large_train_%j.out
|
| 10 |
+
#SBATCH --error=logs/phase_a2_large_train_%j.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase A2: Train large capacity model on 10K dataset
|
| 16 |
+
# Target: 40%+ policy success (vs current 29.67%)
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_a2_large_model"
|
| 25 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 26 |
+
|
| 27 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 28 |
+
|
| 29 |
+
echo "=== Phase A2: Training Large Capacity Model ==="
|
| 30 |
+
echo "Seed: $SEED"
|
| 31 |
+
echo "Hidden dim: 512 (vs current 256)"
|
| 32 |
+
echo "Dataset: 10K groups"
|
| 33 |
+
echo "Target: 40%+ policy success"
|
| 34 |
+
echo ""
|
| 35 |
+
|
| 36 |
+
python scripts/train_dovla.py \
|
| 37 |
+
--dataset "$DATASET" \
|
| 38 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 39 |
+
--objective lattice_field \
|
| 40 |
+
--hidden-dim 512 \
|
| 41 |
+
--action-horizon 4 \
|
| 42 |
+
--epochs 100 \
|
| 43 |
+
--batch-groups 16 \
|
| 44 |
+
--records-per-group 8 \
|
| 45 |
+
--lr 0.0003 \
|
| 46 |
+
--weight-decay 0.01 \
|
| 47 |
+
--device auto \
|
| 48 |
+
--seed $SEED \
|
| 49 |
+
--observation-mode state \
|
| 50 |
+
--loss-weight bc=1.0 \
|
| 51 |
+
--loss-weight field_effect=1.0 \
|
| 52 |
+
--loss-weight field_potential=1.0 \
|
| 53 |
+
--loss-weight field_preference=0.5 \
|
| 54 |
+
--loss-weight field_anchor=0.1
|
| 55 |
+
|
| 56 |
+
echo ""
|
| 57 |
+
echo "✅ Phase A2 complete: Large model trained (seed $SEED)"
|
| 58 |
+
echo ""
|
| 59 |
+
echo "Next: Run phase_a3_eval_large_model.sbatch"
|
scripts/slurm/phase_a3_eval_large_model.sbatch
ADDED
|
@@ -0,0 +1,50 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_large_eval
|
| 3 |
+
#SBATCH --partition=${DOVLA_PARTITION:-compute}
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:1
|
| 8 |
+
#SBATCH --mem=32G
|
| 9 |
+
#SBATCH --time=12:00:00
|
| 10 |
+
#SBATCH --output=logs/phase_a3_large_eval_%j.out
|
| 11 |
+
#SBATCH --error=logs/phase_a3_large_eval_%j.err
|
| 12 |
+
#SBATCH --array=0-2
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
# Phase A3: Evaluate large model with lattice eval + policy rollout
|
| 17 |
+
# Target: Measure improvement over baseline 29.67%
|
| 18 |
+
|
| 19 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 20 |
+
cd "$PROJECT_DIR"
|
| 21 |
+
|
| 22 |
+
source .venv/bin/activate
|
| 23 |
+
|
| 24 |
+
CHECKPOINT_DIR="/scratch/$USER/dovla/experiments/phase_a2_large_model/seed_$SLURM_ARRAY_TASK_ID"
|
| 25 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 26 |
+
OUT_DIR="$CHECKPOINT_DIR"
|
| 27 |
+
|
| 28 |
+
echo "=== Phase A3: Evaluating Large Model (seed $SLURM_ARRAY_TASK_ID) ==="
|
| 29 |
+
echo ""
|
| 30 |
+
|
| 31 |
+
# Lattice evaluation
|
| 32 |
+
echo "Running lattice evaluation..."
|
| 33 |
+
python scripts/eval_lattice_checkpoint.py \
|
| 34 |
+
--checkpoint "$CHECKPOINT_DIR/best.pt" \
|
| 35 |
+
--dataset "$DATASET" \
|
| 36 |
+
--out "$OUT_DIR/lattice_eval.json" \
|
| 37 |
+
--mode field_only \
|
| 38 |
+
--all-groups
|
| 39 |
+
|
| 40 |
+
echo ""
|
| 41 |
+
echo "Running policy rollout..."
|
| 42 |
+
python scripts/eval_maniskill_policy_rollout.py \
|
| 43 |
+
--checkpoint "$CHECKPOINT_DIR/best.pt" \
|
| 44 |
+
--dataset "$DATASET" \
|
| 45 |
+
--out "$OUT_DIR/policy_rollout.json" \
|
| 46 |
+
--num-groups 700 \
|
| 47 |
+
--mode validation
|
| 48 |
+
|
| 49 |
+
echo ""
|
| 50 |
+
echo "✅ Phase A3 complete: Evaluation done (seed $SLURM_ARRAY_TASK_ID)"
|
scripts/slurm/phase_a4_hparam_sweep.sbatch
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_hparam_sweep
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_a4_hparam_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/phase_a4_hparam_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-8
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase A4: Hyperparameter sweep
|
| 16 |
+
# Grid: 3 LR x 3 hidden_dim = 9 configs
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_ROOT="/scratch/$USER/dovla/experiments/phase_a4_hparam_sweep"
|
| 25 |
+
|
| 26 |
+
# Hyperparameter grid
|
| 27 |
+
LRS=(0.0001 0.0003 0.001)
|
| 28 |
+
HIDDEN_DIMS=(256 512 1024)
|
| 29 |
+
|
| 30 |
+
# Map array index to config
|
| 31 |
+
IDX=$SLURM_ARRAY_TASK_ID
|
| 32 |
+
LR_IDX=$((IDX / 3))
|
| 33 |
+
HD_IDX=$((IDX % 3))
|
| 34 |
+
|
| 35 |
+
LR="${LRS[$LR_IDX]}"
|
| 36 |
+
HIDDEN_DIM="${HIDDEN_DIMS[$HD_IDX]}"
|
| 37 |
+
|
| 38 |
+
OUT_DIR="$OUT_ROOT/lr${LR}_h${HIDDEN_DIM}"
|
| 39 |
+
mkdir -p "$OUT_DIR" logs
|
| 40 |
+
|
| 41 |
+
echo "=== Phase A4: Hyperparameter Sweep ==="
|
| 42 |
+
echo "Config $IDX: LR=$LR, Hidden=$HIDDEN_DIM"
|
| 43 |
+
echo ""
|
| 44 |
+
|
| 45 |
+
python scripts/train_dovla.py \
|
| 46 |
+
--dataset "$DATASET" \
|
| 47 |
+
--out "$OUT_DIR" \
|
| 48 |
+
--objective lattice_field \
|
| 49 |
+
--hidden-dim "$HIDDEN_DIM" \
|
| 50 |
+
--epochs 50 \
|
| 51 |
+
--batch-groups 16 \
|
| 52 |
+
--lr "$LR" \
|
| 53 |
+
--device auto \
|
| 54 |
+
--seed 0
|
| 55 |
+
|
| 56 |
+
echo ""
|
| 57 |
+
# Quick eval
|
| 58 |
+
python scripts/eval_lattice_checkpoint.py \
|
| 59 |
+
--checkpoint "$OUT_DIR/best.pt" \
|
| 60 |
+
--dataset /scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection \
|
| 61 |
+
--out "$OUT_DIR/lattice_eval.json" \
|
| 62 |
+
--mode field_only \
|
| 63 |
+
--all-groups
|
| 64 |
+
|
| 65 |
+
echo "✅ Phase A4 config $IDX complete"
|
scripts/slurm/phase_a5_horizon_sweep.sbatch
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_horizon_sweep
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=48000M
|
| 8 |
+
#SBATCH --time=24:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_a5_horizon_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/phase_a5_horizon_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-3
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase A5: Action horizon sweep
|
| 16 |
+
# Test H=4,8,12,16 to see if longer horizons help
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_ROOT="/scratch/$USER/dovla/experiments/phase_a5_horizon_sweep"
|
| 25 |
+
|
| 26 |
+
HORIZONS=(4 8 12 16)
|
| 27 |
+
HORIZON="${HORIZONS[$SLURM_ARRAY_TASK_ID]}"
|
| 28 |
+
|
| 29 |
+
OUT_DIR="$OUT_ROOT/h${HORIZON}"
|
| 30 |
+
mkdir -p "$OUT_DIR" logs
|
| 31 |
+
|
| 32 |
+
echo "=== Phase A5: Action Horizon Sweep ==="
|
| 33 |
+
echo "Horizon: $HORIZON (current baseline: 4)"
|
| 34 |
+
echo ""
|
| 35 |
+
|
| 36 |
+
python scripts/train_dovla.py \
|
| 37 |
+
--dataset "$DATASET" \
|
| 38 |
+
--out "$OUT_DIR" \
|
| 39 |
+
--objective lattice_field \
|
| 40 |
+
--hidden-dim 512 \
|
| 41 |
+
--action-horizon "$HORIZON" \
|
| 42 |
+
--epochs 50 \
|
| 43 |
+
--batch-groups 16 \
|
| 44 |
+
--lr 0.0003 \
|
| 45 |
+
--device auto \
|
| 46 |
+
--seed 0
|
| 47 |
+
|
| 48 |
+
echo ""
|
| 49 |
+
python scripts/eval_lattice_checkpoint.py \
|
| 50 |
+
--checkpoint "$OUT_DIR/best.pt" \
|
| 51 |
+
--dataset /scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection \
|
| 52 |
+
--out "$OUT_DIR/lattice_eval.json" \
|
| 53 |
+
--mode field_only \
|
| 54 |
+
--all-groups
|
| 55 |
+
|
| 56 |
+
python scripts/eval_maniskill_policy_rollout.py \
|
| 57 |
+
--checkpoint "$OUT_DIR/best.pt" \
|
| 58 |
+
--dataset /scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection \
|
| 59 |
+
--out "$OUT_DIR/policy_rollout.json" \
|
| 60 |
+
--num-groups 700 \
|
| 61 |
+
--mode validation
|
| 62 |
+
|
| 63 |
+
echo "✅ Phase A5 horizon=$HORIZON complete"
|
scripts/slurm/phase_b_generate_12tasks.sbatch
ADDED
|
@@ -0,0 +1,109 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_12task_gen
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=16
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=72:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_b_12task_gen_%j.out
|
| 10 |
+
#SBATCH --error=logs/phase_b_12task_gen_%j.err
|
| 11 |
+
|
| 12 |
+
set -euo pipefail
|
| 13 |
+
|
| 14 |
+
# Phase B Option 1: Generate 12-task ManiSkill collection
|
| 15 |
+
# Fastest option - uses existing infrastructure
|
| 16 |
+
|
| 17 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 18 |
+
cd "$PROJECT_DIR"
|
| 19 |
+
|
| 20 |
+
source .venv/bin/activate
|
| 21 |
+
|
| 22 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_b_12task_collection"
|
| 23 |
+
K=16
|
| 24 |
+
STATE_BATCH_SIZE=16
|
| 25 |
+
|
| 26 |
+
# 12 tasks: 6 existing + 6 new
|
| 27 |
+
declare -A TASK_GROUPS=(
|
| 28 |
+
# Original 6
|
| 29 |
+
["PickCube-v1"]=800
|
| 30 |
+
["PushCube-v1"]=800
|
| 31 |
+
["PullCube-v1"]=600
|
| 32 |
+
["StackCube-v1"]=600
|
| 33 |
+
["LiftPegUpright-v1"]=600
|
| 34 |
+
["PegInsertionSide-v1"]=600
|
| 35 |
+
|
| 36 |
+
# New 6 (TODO: ensure demos exist)
|
| 37 |
+
["TurnFaucet-v1"]=500
|
| 38 |
+
["OpenDrawer-v1"]=500
|
| 39 |
+
["CloseDrawer-v1"]=500
|
| 40 |
+
["PlugCharger-v1"]=400
|
| 41 |
+
["HangMug-v1"]=400
|
| 42 |
+
["PourWater-v1"]=400
|
| 43 |
+
)
|
| 44 |
+
|
| 45 |
+
mkdir -p "$OUT_DIR" logs
|
| 46 |
+
|
| 47 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 48 |
+
echo "Phase B Option 1: 12-Task ManiSkill Collection"
|
| 49 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 50 |
+
echo ""
|
| 51 |
+
echo "Target: 6,200 groups, 99,200 records (K=$K)"
|
| 52 |
+
echo "Strategy: Expand existing ManiSkill tasks"
|
| 53 |
+
echo ""
|
| 54 |
+
|
| 55 |
+
# Check if this is just a planning run
|
| 56 |
+
if [ "${DRY_RUN:-0}" = "1" ]; then
|
| 57 |
+
echo "DRY RUN: Would generate 12 tasks"
|
| 58 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 59 |
+
echo " $TASK: ${TASK_GROUPS[$TASK]} groups"
|
| 60 |
+
done
|
| 61 |
+
exit 0
|
| 62 |
+
fi
|
| 63 |
+
|
| 64 |
+
# Generate each task
|
| 65 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 66 |
+
NUM_GROUPS="${TASK_GROUPS[$TASK]}"
|
| 67 |
+
|
| 68 |
+
# Check if already generated
|
| 69 |
+
if [ -d "$OUT_DIR/${TASK}_k${K}_n${NUM_GROUPS}/merged" ]; then
|
| 70 |
+
echo "✓ $TASK already exists, skipping"
|
| 71 |
+
continue
|
| 72 |
+
fi
|
| 73 |
+
|
| 74 |
+
echo "Generating $TASK: $NUM_GROUPS groups..."
|
| 75 |
+
|
| 76 |
+
# Use existing generation script
|
| 77 |
+
python scripts/generate_maniskill_lattice.py \
|
| 78 |
+
--env-id "$TASK" \
|
| 79 |
+
--control-mode pd_ee_delta_pose \
|
| 80 |
+
--out "$OUT_DIR/${TASK}_k${K}_n${NUM_GROUPS}" \
|
| 81 |
+
--num-groups "$NUM_GROUPS" \
|
| 82 |
+
--k "$K" \
|
| 83 |
+
--state-batch-size "$STATE_BATCH_SIZE" \
|
| 84 |
+
--seed 42 \
|
| 85 |
+
--pre-success-only \
|
| 86 |
+
--use-official-demos || {
|
| 87 |
+
echo "⚠️ $TASK failed (demo might not exist)"
|
| 88 |
+
continue
|
| 89 |
+
}
|
| 90 |
+
|
| 91 |
+
echo "✅ $TASK complete"
|
| 92 |
+
echo ""
|
| 93 |
+
done
|
| 94 |
+
|
| 95 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 96 |
+
echo "Merging into unified 12-task collection"
|
| 97 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 98 |
+
|
| 99 |
+
python scripts/make_cil_collection.py \
|
| 100 |
+
--source-dirs "$OUT_DIR"/*/merged \
|
| 101 |
+
--out "$OUT_DIR/merged_12tasks" \
|
| 102 |
+
--name "phase_b_12task_collection"
|
| 103 |
+
|
| 104 |
+
echo ""
|
| 105 |
+
echo "✅ Phase B Option 1 complete: 12-task collection ready"
|
| 106 |
+
echo " Location: $OUT_DIR/merged_12tasks"
|
| 107 |
+
echo ""
|
| 108 |
+
echo "Next: Train on 12 tasks"
|
| 109 |
+
echo " sbatch scripts/slurm/phase_b_train_12tasks.sbatch"
|
scripts/slurm/phase_b_train_12tasks.sbatch
ADDED
|
@@ -0,0 +1,64 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_12task_train
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=96:00:00
|
| 9 |
+
#SBATCH --output=logs/phase_b_12task_train_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/phase_b_12task_train_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# Phase B: Train on 12-task collection (3 seeds)
|
| 16 |
+
|
| 17 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 18 |
+
cd "$PROJECT_DIR"
|
| 19 |
+
|
| 20 |
+
source .venv/bin/activate
|
| 21 |
+
|
| 22 |
+
DATASET="/scratch/$USER/dovla/experiments/phase_b_12task_collection/merged_12tasks"
|
| 23 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/phase_b_12task_model"
|
| 24 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 25 |
+
|
| 26 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 27 |
+
|
| 28 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 29 |
+
echo "Phase B: Training on 12-Task Collection"
|
| 30 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 31 |
+
echo ""
|
| 32 |
+
echo "Seed: $SEED"
|
| 33 |
+
echo "Tasks: 12 (6 existing + 6 new)"
|
| 34 |
+
echo "Groups: ~6,200"
|
| 35 |
+
echo "Hidden dim: 1024 (larger for 12 tasks)"
|
| 36 |
+
echo ""
|
| 37 |
+
|
| 38 |
+
python scripts/train_dovla.py \
|
| 39 |
+
--dataset "$DATASET" \
|
| 40 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 41 |
+
--objective lattice_field \
|
| 42 |
+
--hidden-dim 1024 \
|
| 43 |
+
--action-horizon 4 \
|
| 44 |
+
--epochs 100 \
|
| 45 |
+
--batch-groups 16 \
|
| 46 |
+
--records-per-group 8 \
|
| 47 |
+
--lr 0.0003 \
|
| 48 |
+
--weight-decay 0.01 \
|
| 49 |
+
--dropout 0.1 \
|
| 50 |
+
--warmup-steps 1000 \
|
| 51 |
+
--device auto \
|
| 52 |
+
--seed $SEED \
|
| 53 |
+
--observation-mode state \
|
| 54 |
+
--loss-weight bc=1.0 \
|
| 55 |
+
--loss-weight field_effect=1.0 \
|
| 56 |
+
--loss-weight field_utility_regression=1.0 \
|
| 57 |
+
--loss-weight field_utility_margin=0.5 \
|
| 58 |
+
--loss-weight field_preference=0.5 \
|
| 59 |
+
--loss-weight effect_anchor=0.1
|
| 60 |
+
|
| 61 |
+
echo ""
|
| 62 |
+
echo "✅ Phase B training complete (seed $SEED)"
|
| 63 |
+
echo ""
|
| 64 |
+
echo "Next: Evaluate on held-out tasks"
|
scripts/slurm/plan_c_generate_10k.sbatch
ADDED
|
@@ -0,0 +1,121 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_10k_planc
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=16
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=96:00:00
|
| 9 |
+
#SBATCH --output=logs/plan_c_10k_gen_%j.out
|
| 10 |
+
#SBATCH --error=logs/plan_c_10k_gen_%j.err
|
| 11 |
+
|
| 12 |
+
set -euo pipefail
|
| 13 |
+
|
| 14 |
+
# Plan C: Phase 1B - Generate 10K Groups with Enhanced Sampling
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 17 |
+
cd "$PROJECT_DIR"
|
| 18 |
+
|
| 19 |
+
source .venv/bin/activate
|
| 20 |
+
|
| 21 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/plan_c_10k_enhanced"
|
| 22 |
+
K=16
|
| 23 |
+
STATE_BATCH_SIZE=16
|
| 24 |
+
DEMO_BASE="/scratch/$USER/dovla/maniskill_multitask_demos"
|
| 25 |
+
|
| 26 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 27 |
+
echo "Plan C: 10K Generation with Enhanced Sampling"
|
| 28 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 29 |
+
echo ""
|
| 30 |
+
echo "Target: 48-50%+ policy success"
|
| 31 |
+
echo "Strategy: Maximum quality improvements"
|
| 32 |
+
echo ""
|
| 33 |
+
|
| 34 |
+
# Task distribution (balanced across difficulty)
|
| 35 |
+
declare -A TASK_GROUPS=(
|
| 36 |
+
["PickCube-v1"]=1800
|
| 37 |
+
["PushCube-v1"]=1800
|
| 38 |
+
["PullCube-v1"]=1600
|
| 39 |
+
["StackCube-v1"]=1600
|
| 40 |
+
["LiftPegUpright-v1"]=1600
|
| 41 |
+
["PegInsertionSide-v1"]=1600
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
declare -A TASK_DEMOS=(
|
| 45 |
+
["PickCube-v1"]="$DEMO_BASE/PickCube-v1/motionplanning/trajectory.h5"
|
| 46 |
+
["PushCube-v1"]="$DEMO_BASE/PushCube-v1/motionplanning/trajectory.h5"
|
| 47 |
+
["PullCube-v1"]="$DEMO_BASE/PullCube-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 48 |
+
["StackCube-v1"]="$DEMO_BASE/StackCube-v1/motionplanning/trajectory.h5"
|
| 49 |
+
["LiftPegUpright-v1"]="$DEMO_BASE/LiftPegUpright-v1/rl/trajectory.none.pd_ee_delta_pose.physx_cuda.h5"
|
| 50 |
+
["PegInsertionSide-v1"]="$DEMO_BASE/PegInsertionSide-v1/motionplanning/trajectory.h5"
|
| 51 |
+
)
|
| 52 |
+
|
| 53 |
+
TOTAL_GROUPS=0
|
| 54 |
+
for count in "${TASK_GROUPS[@]}"; do
|
| 55 |
+
TOTAL_GROUPS=$((TOTAL_GROUPS + count))
|
| 56 |
+
done
|
| 57 |
+
|
| 58 |
+
echo "Task distribution (total: $TOTAL_GROUPS groups):"
|
| 59 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 60 |
+
echo " ${TASK}: ${TASK_GROUPS[$TASK]} groups"
|
| 61 |
+
done
|
| 62 |
+
echo ""
|
| 63 |
+
|
| 64 |
+
# Generate each task
|
| 65 |
+
for TASK in "${!TASK_GROUPS[@]}"; do
|
| 66 |
+
NUM_GROUPS="${TASK_GROUPS[$TASK]}"
|
| 67 |
+
DEMO_PATH="${TASK_DEMOS[$TASK]}"
|
| 68 |
+
TASK_OUT="$OUT_DIR/${TASK}_k${K}_n${NUM_GROUPS}"
|
| 69 |
+
|
| 70 |
+
if [ -d "$TASK_OUT/merged" ]; then
|
| 71 |
+
echo "✓ $TASK already generated, skipping"
|
| 72 |
+
continue
|
| 73 |
+
fi
|
| 74 |
+
|
| 75 |
+
if [ ! -f "$DEMO_PATH" ]; then
|
| 76 |
+
echo "❌ Demo not found: $DEMO_PATH"
|
| 77 |
+
echo " Trying alternate location..."
|
| 78 |
+
# Try RL demos as fallback
|
| 79 |
+
DEMO_PATH="$DEMO_BASE/${TASK}/rl/trajectory.h5"
|
| 80 |
+
if [ ! -f "$DEMO_PATH" ]; then
|
| 81 |
+
echo " ❌ No demo found, skipping $TASK"
|
| 82 |
+
continue
|
| 83 |
+
fi
|
| 84 |
+
fi
|
| 85 |
+
|
| 86 |
+
echo "Generating $TASK: $NUM_GROUPS groups..."
|
| 87 |
+
echo " Demo: $DEMO_PATH"
|
| 88 |
+
echo " Start: $(date)"
|
| 89 |
+
|
| 90 |
+
python scripts/generate_maniskill_lattice.py \
|
| 91 |
+
--demo "$DEMO_PATH" \
|
| 92 |
+
--env-id "$TASK" \
|
| 93 |
+
--control-mode pd_ee_delta_pose \
|
| 94 |
+
--out "$TASK_OUT" \
|
| 95 |
+
--num-groups "$NUM_GROUPS" \
|
| 96 |
+
--k "$K" \
|
| 97 |
+
--state-batch-size "$STATE_BATCH_SIZE" \
|
| 98 |
+
--seed 42 \
|
| 99 |
+
--candidate-mode structured
|
| 100 |
+
|
| 101 |
+
echo " ✅ Complete: $(date)"
|
| 102 |
+
echo ""
|
| 103 |
+
done
|
| 104 |
+
|
| 105 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 106 |
+
echo "Merging all tasks into unified collection"
|
| 107 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 108 |
+
|
| 109 |
+
python scripts/make_cil_collection.py \
|
| 110 |
+
--source-dirs "$OUT_DIR"/*/merged \
|
| 111 |
+
--out "$OUT_DIR/merged_10k" \
|
| 112 |
+
--name "plan_c_10k_enhanced"
|
| 113 |
+
|
| 114 |
+
echo ""
|
| 115 |
+
echo "✅ Plan C Phase 1B Complete!"
|
| 116 |
+
echo ""
|
| 117 |
+
echo "Output: $OUT_DIR/merged_10k"
|
| 118 |
+
echo "Total groups: $TOTAL_GROUPS"
|
| 119 |
+
echo "Total records: $((TOTAL_GROUPS * K))"
|
| 120 |
+
echo ""
|
| 121 |
+
echo "Next: Phase 2A - Attention architecture"
|
scripts/slurm/prepare_maniskill_baselines.sbatch
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_baseprep
|
| 3 |
+
#SBATCH --account=def-yalda_cpu
|
| 4 |
+
#SBATCH --partition=cpubase_bycore_b1
|
| 5 |
+
#SBATCH --nodes=1
|
| 6 |
+
#SBATCH --ntasks=1
|
| 7 |
+
#SBATCH --cpus-per-task=4
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=01:00:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
DATASET="${DATASET:?Set DATASET to the measured CIL collection}"
|
| 17 |
+
OUT_ROOT="${OUT_ROOT:?Set OUT_ROOT for transformed datasets}"
|
| 18 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 19 |
+
|
| 20 |
+
cd "$PROJECT_DIR"
|
| 21 |
+
mkdir -p "$OUT_ROOT"
|
| 22 |
+
|
| 23 |
+
"$PYTHON" scripts/prepare_baseline_dataset.py \
|
| 24 |
+
--dataset "$DATASET" \
|
| 25 |
+
--baseline expert_only_bc \
|
| 26 |
+
--out "$OUT_ROOT/expert_only_bc" \
|
| 27 |
+
--shard-size 2048
|
| 28 |
+
|
| 29 |
+
"$PYTHON" scripts/prepare_baseline_dataset.py \
|
| 30 |
+
--dataset "$DATASET" \
|
| 31 |
+
--baseline label_only_counterfactual \
|
| 32 |
+
--out "$OUT_ROOT/label_only_counterfactual" \
|
| 33 |
+
--shard-size 2048
|
scripts/slurm/render_maniskill_multitask.sbatch
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_multi_rgb
|
| 3 |
+
#SBATCH --account=def-yalda_cpu
|
| 4 |
+
#SBATCH --partition=cpubase_bycore_b1
|
| 5 |
+
#SBATCH --nodes=1
|
| 6 |
+
#SBATCH --ntasks=1
|
| 7 |
+
#SBATCH --cpus-per-task=8
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=03:00:00
|
| 10 |
+
#SBATCH --array=0-4%2
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
MULTITASK_OUT_ROOT="${MULTITASK_OUT_ROOT:?Set MULTITASK_OUT_ROOT}"
|
| 18 |
+
|
| 19 |
+
case "${SLURM_ARRAY_TASK_ID:-0}" in
|
| 20 |
+
0) ENV_ID="PushCube-v1" ;;
|
| 21 |
+
1) ENV_ID="PullCube-v1" ;;
|
| 22 |
+
2) ENV_ID="StackCube-v1" ;;
|
| 23 |
+
3) ENV_ID="LiftPegUpright-v1" ;;
|
| 24 |
+
4) ENV_ID="PegInsertionSide-v1" ;;
|
| 25 |
+
*) echo "unsupported array index" >&2; exit 2 ;;
|
| 26 |
+
esac
|
| 27 |
+
|
| 28 |
+
export PROJECT_DIR
|
| 29 |
+
export DATASET="$MULTITASK_OUT_ROOT/$ENV_ID"
|
| 30 |
+
export IMAGE_QUALITY="${IMAGE_QUALITY:-85}"
|
| 31 |
+
export SEED="${SEED:-0}"
|
| 32 |
+
|
| 33 |
+
exec bash "$PROJECT_DIR/scripts/slurm/render_maniskill_observations.sbatch"
|
scripts/slurm/render_maniskill_observations.sbatch
ADDED
|
@@ -0,0 +1,49 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_render
|
| 3 |
+
#SBATCH --account=def-yalda_cpu
|
| 4 |
+
#SBATCH --partition=cpubase_bycore_b1
|
| 5 |
+
#SBATCH --nodes=1
|
| 6 |
+
#SBATCH --ntasks=1
|
| 7 |
+
#SBATCH --cpus-per-task=8
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=02:00:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 17 |
+
SIF="$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif"
|
| 18 |
+
PYTHON="$SCRATCH_ROOT/envs/maniskill/bin/python"
|
| 19 |
+
NATIVE_LIBS="$SCRATCH_ROOT/native_libs/lib"
|
| 20 |
+
CPU_RENDER_LIBS="$SCRATCH_ROOT/cpu_render_libs"
|
| 21 |
+
VULKAN_ICD="$CPU_RENDER_LIBS/share/vulkan/icd.d/lvp_icd.x86_64.json"
|
| 22 |
+
|
| 23 |
+
DATASET="${DATASET:?Set DATASET to a generated ManiSkill CIL directory}"
|
| 24 |
+
IMAGE_QUALITY="${IMAGE_QUALITY:-90}"
|
| 25 |
+
SEED="${SEED:-0}"
|
| 26 |
+
OVERWRITE="${OVERWRITE:-0}"
|
| 27 |
+
RUNTIME_DIR="/tmp/$USER/dovla-render-$SLURM_JOB_ID"
|
| 28 |
+
CACHE_DIR="/tmp/$USER/dovla-render-cache-$SLURM_JOB_ID"
|
| 29 |
+
|
| 30 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 31 |
+
cd "$PROJECT_DIR"
|
| 32 |
+
mkdir -p "$RUNTIME_DIR" "$CACHE_DIR"
|
| 33 |
+
chmod 700 "$RUNTIME_DIR"
|
| 34 |
+
|
| 35 |
+
ARGS=(
|
| 36 |
+
--dataset "$DATASET"
|
| 37 |
+
--render-backend cpu
|
| 38 |
+
--image-quality "$IMAGE_QUALITY"
|
| 39 |
+
--seed "$SEED"
|
| 40 |
+
)
|
| 41 |
+
if [[ "$OVERWRITE" == "1" ]]; then
|
| 42 |
+
ARGS+=(--overwrite)
|
| 43 |
+
fi
|
| 44 |
+
|
| 45 |
+
apptainer exec \
|
| 46 |
+
--env "LD_LIBRARY_PATH=$CPU_RENDER_LIBS/lib:$NATIVE_LIBS,XDG_RUNTIME_DIR=$RUNTIME_DIR,MESA_SHADER_CACHE_DIR=$CACHE_DIR,LIBGL_ALWAYS_SOFTWARE=1,LP_NUM_THREADS=4,VK_ICD_FILENAMES=$VULKAN_ICD,VK_DRIVER_FILES=$VULKAN_ICD,OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1" \
|
| 47 |
+
"$SIF" "$PYTHON" scripts/render_maniskill_observations.py "${ARGS[@]}"
|
| 48 |
+
|
| 49 |
+
rm -rf "$RUNTIME_DIR" "$CACHE_DIR"
|
scripts/slurm/run_external_vla_baseline.sbatch
ADDED
|
@@ -0,0 +1,51 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ext_vla
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=40G
|
| 9 |
+
#SBATCH --time=04:00:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
MODEL_FAMILY="${MODEL_FAMILY:-smolvla}"
|
| 17 |
+
DATASET="${DATASET:?Set DATASET to the held-out DoVLA-CIL dataset}"
|
| 18 |
+
OUT="${OUT:?Set OUT to a run directory}"
|
| 19 |
+
CHECKPOINT="${CHECKPOINT:-}"
|
| 20 |
+
ADAPTER_ENTRYPOINT="${ADAPTER_ENTRYPOINT:-}"
|
| 21 |
+
ADAPTER_CONFIG="${ADAPTER_CONFIG:-}"
|
| 22 |
+
PYTHON="${PYTHON:-python}"
|
| 23 |
+
DRY_RUN="${DRY_RUN:-0}"
|
| 24 |
+
|
| 25 |
+
cd "$PROJECT_DIR"
|
| 26 |
+
mkdir -p "$OUT" outputs/hpc/logs
|
| 27 |
+
|
| 28 |
+
ARGS=(
|
| 29 |
+
scripts/run_external_vla_baseline.py
|
| 30 |
+
--model-family "$MODEL_FAMILY"
|
| 31 |
+
--dataset "$DATASET"
|
| 32 |
+
--out "$OUT"
|
| 33 |
+
--python "$PYTHON"
|
| 34 |
+
)
|
| 35 |
+
|
| 36 |
+
if [[ -n "$CHECKPOINT" ]]; then
|
| 37 |
+
ARGS+=(--checkpoint "$CHECKPOINT")
|
| 38 |
+
fi
|
| 39 |
+
if [[ -n "$ADAPTER_ENTRYPOINT" ]]; then
|
| 40 |
+
ARGS+=(--adapter-entrypoint "$ADAPTER_ENTRYPOINT")
|
| 41 |
+
fi
|
| 42 |
+
if [[ -n "$ADAPTER_CONFIG" ]]; then
|
| 43 |
+
ARGS+=(--adapter-config "$ADAPTER_CONFIG")
|
| 44 |
+
fi
|
| 45 |
+
if [[ "$DRY_RUN" == "1" ]]; then
|
| 46 |
+
ARGS+=(--dry-run)
|
| 47 |
+
else
|
| 48 |
+
ARGS+=(--require-ready)
|
| 49 |
+
fi
|
| 50 |
+
|
| 51 |
+
"$PYTHON" "${ARGS[@]}"
|
scripts/slurm/run_scaling.sbatch
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=${DOVLA_JOB_NAME:-dovla_scaling}
|
| 3 |
+
#SBATCH --partition=${DOVLA_PARTITION:-gpu}
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=${DOVLA_CPUS_PER_TASK:-8}
|
| 7 |
+
#SBATCH --gres=gpu:${DOVLA_GPUS_PER_TASK:-1}
|
| 8 |
+
#SBATCH --mem=${DOVLA_MEM:-64G}
|
| 9 |
+
#SBATCH --time=${DOVLA_TIME:-24:00:00}
|
| 10 |
+
#SBATCH --output=${DOVLA_LOG_DIR:-logs/slurm}/%x_%j.out
|
| 11 |
+
#SBATCH --error=${DOVLA_LOG_DIR:-logs/slurm}/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 16 |
+
VENV_PATH="${VENV_PATH:-$PROJECT_DIR/.venv}"
|
| 17 |
+
BACKEND="${BACKEND:-toy}"
|
| 18 |
+
TASKS="${TASKS:-builtins}"
|
| 19 |
+
OUT_DIR="${OUT_DIR:-$PROJECT_DIR/runs/scaling_toy}"
|
| 20 |
+
TOTAL_RECORDS="${TOTAL_RECORDS:-4096}"
|
| 21 |
+
K_VALUES="${K_VALUES:-1,2,4,8,16,32}"
|
| 22 |
+
EPOCHS="${EPOCHS:-3}"
|
| 23 |
+
SEED="${SEED:-0}"
|
| 24 |
+
DEVICE="${DEVICE:-auto}"
|
| 25 |
+
|
| 26 |
+
mkdir -p "${DOVLA_LOG_DIR:-logs/slurm}" "$OUT_DIR"
|
| 27 |
+
cd "$PROJECT_DIR"
|
| 28 |
+
|
| 29 |
+
if [ -f "$VENV_PATH/bin/activate" ]; then
|
| 30 |
+
# shellcheck disable=SC1091
|
| 31 |
+
source "$VENV_PATH/bin/activate"
|
| 32 |
+
fi
|
| 33 |
+
|
| 34 |
+
python scripts/run_scaling.py \
|
| 35 |
+
--backend "$BACKEND" \
|
| 36 |
+
--tasks "$TASKS" \
|
| 37 |
+
--out "$OUT_DIR" \
|
| 38 |
+
--total-records "$TOTAL_RECORDS" \
|
| 39 |
+
--k-values "$K_VALUES" \
|
| 40 |
+
--epochs "$EPOCHS" \
|
| 41 |
+
--seed "$SEED" \
|
| 42 |
+
--device "$DEVICE"
|
scripts/slurm/run_smolvla_cil_baseline.sbatch
ADDED
|
@@ -0,0 +1,48 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_smolvla_cil
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_3g.40gb:1
|
| 8 |
+
#SBATCH --mem=40G
|
| 9 |
+
#SBATCH --time=02:00:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
SCRATCH_ROOT="${SCRATCH_ROOT:-/scratch/$USER/dovla}"
|
| 17 |
+
CONTAINER="${CONTAINER:-$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif}"
|
| 18 |
+
PYTHON="${PYTHON:-$SCRATCH_ROOT/envs/smolvla/bin/python}"
|
| 19 |
+
CHECKPOINT="${CHECKPOINT:-$SCRATCH_ROOT/models/smolvla_base-c83c316}"
|
| 20 |
+
DATASET="${DATASET:-$SCRATCH_ROOT/experiments/maniskill_presuccess_six_task_collection}"
|
| 21 |
+
ADAPTER_CONFIG="${ADAPTER_CONFIG:-$PROJECT_DIR/configs/external/smolvla_cil_smoke.json}"
|
| 22 |
+
OUT="${OUT:-$SCRATCH_ROOT/experiments/smolvla_cil_smoke}"
|
| 23 |
+
CONTAINER_ADAPTER_CONFIG="$ADAPTER_CONFIG"
|
| 24 |
+
if [[ "$ADAPTER_CONFIG" == "$PROJECT_DIR/"* ]]; then
|
| 25 |
+
CONTAINER_ADAPTER_CONFIG="/workspace/${ADAPTER_CONFIG#"$PROJECT_DIR/"}"
|
| 26 |
+
fi
|
| 27 |
+
|
| 28 |
+
cd "$PROJECT_DIR"
|
| 29 |
+
mkdir -p "$OUT" outputs/hpc/logs
|
| 30 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 31 |
+
|
| 32 |
+
apptainer exec \
|
| 33 |
+
--nv \
|
| 34 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 35 |
+
-B "$PROJECT_DIR:/workspace" \
|
| 36 |
+
--env \
|
| 37 |
+
"PYTHONNOUSERSITE=1,HF_HUB_OFFLINE=1,TRANSFORMERS_OFFLINE=1,SCRATCH_ROOT=$SCRATCH_ROOT" \
|
| 38 |
+
"$CONTAINER" \
|
| 39 |
+
"$PYTHON" /workspace/scripts/run_external_vla_baseline.py \
|
| 40 |
+
--model-family smolvla \
|
| 41 |
+
--checkpoint "$CHECKPOINT" \
|
| 42 |
+
--dataset "$DATASET" \
|
| 43 |
+
--out "$OUT" \
|
| 44 |
+
--python "$PYTHON" \
|
| 45 |
+
--adapter-entrypoint \
|
| 46 |
+
dovla_cil.eval.smolvla_cil_baseline:run_smolvla_cil_baseline \
|
| 47 |
+
--adapter-config "$CONTAINER_ADAPTER_CONFIG" \
|
| 48 |
+
--require-ready
|
scripts/slurm/smoke_smolvla_checkpoint.sbatch
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_smolvla_smoke
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=00:30:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
SCRATCH_ROOT="${SCRATCH_ROOT:-/scratch/$USER/dovla}"
|
| 17 |
+
CONTAINER="${CONTAINER:-$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif}"
|
| 18 |
+
PYTHON="${PYTHON:-$SCRATCH_ROOT/envs/smolvla/bin/python}"
|
| 19 |
+
CHECKPOINT="${CHECKPOINT:-$SCRATCH_ROOT/models/smolvla_base-c83c316}"
|
| 20 |
+
VLM_REVISION="${VLM_REVISION:-7b375e1b73b11138ff12fe22c8f2822d8fe03467}"
|
| 21 |
+
VLM_METADATA="${VLM_METADATA:-$SCRATCH_ROOT/models/SmolVLM2-500M-Video-Instruct-metadata-$VLM_REVISION}"
|
| 22 |
+
OUT="${OUT:-$PROJECT_DIR/outputs/external_vla_smolvla_checkpoint_smoke.json}"
|
| 23 |
+
CONTAINER_OUT="$OUT"
|
| 24 |
+
if [[ "$OUT" == "$PROJECT_DIR/"* ]]; then
|
| 25 |
+
CONTAINER_OUT="/workspace/${OUT#"$PROJECT_DIR/"}"
|
| 26 |
+
fi
|
| 27 |
+
|
| 28 |
+
cd "$PROJECT_DIR"
|
| 29 |
+
mkdir -p "$(dirname "$OUT")" outputs/hpc/logs
|
| 30 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 31 |
+
|
| 32 |
+
apptainer exec \
|
| 33 |
+
--nv \
|
| 34 |
+
-B "$SCRATCH_ROOT:$SCRATCH_ROOT" \
|
| 35 |
+
-B "$PROJECT_DIR:/workspace" \
|
| 36 |
+
--env PYTHONNOUSERSITE=1,HF_HUB_OFFLINE=1,TRANSFORMERS_OFFLINE=1 \
|
| 37 |
+
"$CONTAINER" \
|
| 38 |
+
"$PYTHON" /workspace/scripts/smoke_smolvla_checkpoint.py \
|
| 39 |
+
--checkpoint "$CHECKPOINT" \
|
| 40 |
+
--vlm-metadata "$VLM_METADATA" \
|
| 41 |
+
--out "$CONTAINER_OUT" \
|
| 42 |
+
--device cuda
|
scripts/slurm/train_attention_model.sbatch
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_attention
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/attention_train_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/attention_train_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# CVPR-Ready: DoVLA-Attention Architecture
|
| 16 |
+
# Single principled contribution: Transformer attention for action comparison
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/cvpr_attention_model"
|
| 25 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 26 |
+
|
| 27 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 28 |
+
|
| 29 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 30 |
+
echo "CVPR Experiment: DoVLA-Attention Architecture"
|
| 31 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 32 |
+
echo ""
|
| 33 |
+
echo "Method: Transformer-based attention for action comparison"
|
| 34 |
+
echo "Contribution: Cross-attention + Self-attention + Pairwise head"
|
| 35 |
+
echo "Dataset: 3,500 groups (SAME as baseline for fair comparison)"
|
| 36 |
+
echo "Seed: $SEED"
|
| 37 |
+
echo ""
|
| 38 |
+
echo "Expected: 42-44% success (vs 38.43% MLP baseline)"
|
| 39 |
+
echo ""
|
| 40 |
+
|
| 41 |
+
# Train with attention architecture
|
| 42 |
+
python scripts/train_dovla_attention.py \
|
| 43 |
+
--dataset "$DATASET" \
|
| 44 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 45 |
+
--model attention \
|
| 46 |
+
--hidden-dim 256 \
|
| 47 |
+
--n-heads 4 \
|
| 48 |
+
--n-layers 2 \
|
| 49 |
+
--action-horizon 4 \
|
| 50 |
+
--epochs 50 \
|
| 51 |
+
--batch-groups 16 \
|
| 52 |
+
--records-per-group 8 \
|
| 53 |
+
--lr 0.0003 \
|
| 54 |
+
--weight-decay 0.01 \
|
| 55 |
+
--device auto \
|
| 56 |
+
--seed $SEED \
|
| 57 |
+
--observation-mode state
|
| 58 |
+
|
| 59 |
+
echo ""
|
| 60 |
+
echo "✅ DoVLA-Attention training complete (seed $SEED)"
|
| 61 |
+
echo ""
|
| 62 |
+
echo "Next: Evaluate and compare with MLP baseline (38.43%)"
|
scripts/slurm/train_dovla.sbatch
ADDED
|
@@ -0,0 +1,44 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=${DOVLA_JOB_NAME:-dovla_train}
|
| 3 |
+
#SBATCH --partition=${DOVLA_PARTITION:-gpu}
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=${DOVLA_CPUS_PER_TASK:-8}
|
| 7 |
+
#SBATCH --gres=gpu:${DOVLA_GPUS_PER_TASK:-1}
|
| 8 |
+
#SBATCH --mem=${DOVLA_MEM:-64G}
|
| 9 |
+
#SBATCH --time=${DOVLA_TIME:-24:00:00}
|
| 10 |
+
#SBATCH --output=${DOVLA_LOG_DIR:-logs/slurm}/%x_%j.out
|
| 11 |
+
#SBATCH --error=${DOVLA_LOG_DIR:-logs/slurm}/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 16 |
+
VENV_PATH="${VENV_PATH:-$PROJECT_DIR/.venv}"
|
| 17 |
+
DATASET="${DATASET:-$PROJECT_DIR/data/cil_toy}"
|
| 18 |
+
OUT_DIR="${OUT_DIR:-$PROJECT_DIR/runs/dovla_toy}"
|
| 19 |
+
EPOCHS="${EPOCHS:-5}"
|
| 20 |
+
BATCH_GROUPS="${BATCH_GROUPS:-8}"
|
| 21 |
+
RECORDS_PER_GROUP="${RECORDS_PER_GROUP:-8}"
|
| 22 |
+
HIDDEN_DIM="${HIDDEN_DIM:-256}"
|
| 23 |
+
LR="${LR:-0.001}"
|
| 24 |
+
DEVICE="${DEVICE:-auto}"
|
| 25 |
+
SEED="${SEED:-0}"
|
| 26 |
+
|
| 27 |
+
mkdir -p "${DOVLA_LOG_DIR:-logs/slurm}" "$OUT_DIR"
|
| 28 |
+
cd "$PROJECT_DIR"
|
| 29 |
+
|
| 30 |
+
if [ -f "$VENV_PATH/bin/activate" ]; then
|
| 31 |
+
# shellcheck disable=SC1091
|
| 32 |
+
source "$VENV_PATH/bin/activate"
|
| 33 |
+
fi
|
| 34 |
+
|
| 35 |
+
python scripts/train_dovla.py \
|
| 36 |
+
--dataset "$DATASET" \
|
| 37 |
+
--out "$OUT_DIR" \
|
| 38 |
+
--epochs "$EPOCHS" \
|
| 39 |
+
--batch-groups "$BATCH_GROUPS" \
|
| 40 |
+
--records-per-group "$RECORDS_PER_GROUP" \
|
| 41 |
+
--hidden-dim "$HIDDEN_DIM" \
|
| 42 |
+
--lr "$LR" \
|
| 43 |
+
--device "$DEVICE" \
|
| 44 |
+
--seed "$SEED"
|
scripts/slurm/train_enhanced_model.sbatch
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_enhanced
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/enhanced_train_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/enhanced_train_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# DoVLA-Attention-Enhanced: SOTA Architecture for CVPR
|
| 16 |
+
# Hierarchical attention + Graph NN + Contrastive + Task-adaptive
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/cvpr_enhanced_model"
|
| 25 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 26 |
+
|
| 27 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 28 |
+
|
| 29 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 30 |
+
echo "DoVLA-Attention-Enhanced: SOTA Training"
|
| 31 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 32 |
+
echo ""
|
| 33 |
+
echo "Architecture Components:"
|
| 34 |
+
echo " 1. Hierarchical Attention (local + global)"
|
| 35 |
+
echo " 2. Graph Neural Network (action relationships)"
|
| 36 |
+
echo " 3. Contrastive Learning (better embeddings)"
|
| 37 |
+
echo " 4. Task-Adaptive Layers (multi-task)"
|
| 38 |
+
echo ""
|
| 39 |
+
echo "Dataset: 3,500 groups (fair comparison)"
|
| 40 |
+
echo "Seed: $SEED"
|
| 41 |
+
echo ""
|
| 42 |
+
echo "Expected: 44-47% success (vs 38.43% baseline)"
|
| 43 |
+
echo "Improvement: +5.5-8.5%"
|
| 44 |
+
echo ""
|
| 45 |
+
|
| 46 |
+
python scripts/train_dovla_enhanced.py \
|
| 47 |
+
--dataset "$DATASET" \
|
| 48 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 49 |
+
--hidden-dim 256 \
|
| 50 |
+
--n-heads 4 \
|
| 51 |
+
--n-layers 3 \
|
| 52 |
+
--epochs 50 \
|
| 53 |
+
--batch-size 16 \
|
| 54 |
+
--lr 0.0003 \
|
| 55 |
+
--weight-decay 0.01 \
|
| 56 |
+
--contrastive-weight 0.1 \
|
| 57 |
+
--seed $SEED \
|
| 58 |
+
--device auto
|
| 59 |
+
|
| 60 |
+
echo ""
|
| 61 |
+
echo "✅ Enhanced training complete (seed $SEED)"
|
| 62 |
+
echo ""
|
| 63 |
+
echo "Next: Evaluate and compare with baseline"
|
scripts/slurm/train_h16_policy.sbatch
ADDED
|
@@ -0,0 +1,54 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=train_h16_policy
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:1
|
| 8 |
+
#SBATCH --mem=48G
|
| 9 |
+
#SBATCH --time=04:00:00
|
| 10 |
+
#SBATCH --output=logs/train_h16_%A_%a.out
|
| 11 |
+
#SBATCH --error=logs/train_h16_%A_%a.err
|
| 12 |
+
#SBATCH --array=0-2
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
# Train policy on h=16 collection (oracle 94.76%)
|
| 17 |
+
# Expected: val top-1 ~85-90%, online rollout 55-70%+
|
| 18 |
+
|
| 19 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 20 |
+
cd "$PROJECT_DIR"
|
| 21 |
+
|
| 22 |
+
source .venv/bin/activate
|
| 23 |
+
|
| 24 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 25 |
+
DATASET="/scratch/$USER/dovla/experiments/six_task_h16_collection"
|
| 26 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/h16_policy_runs/seed_$SEED"
|
| 27 |
+
|
| 28 |
+
mkdir -p "$OUT_DIR" logs
|
| 29 |
+
|
| 30 |
+
echo "=================================================="
|
| 31 |
+
echo "Training Policy on h=16 Collection"
|
| 32 |
+
echo "Seed: $SEED"
|
| 33 |
+
echo "Dataset: $DATASET"
|
| 34 |
+
echo "Expected oracle: 94.76%"
|
| 35 |
+
echo "Expected val top-1: 85-90%"
|
| 36 |
+
echo "=================================================="
|
| 37 |
+
|
| 38 |
+
python scripts/train_hybrid_direct.py \
|
| 39 |
+
--dataset "$DATASET" \
|
| 40 |
+
--out "$OUT_DIR" \
|
| 41 |
+
--d-model 256 \
|
| 42 |
+
--n-heads 8 \
|
| 43 |
+
--n-layers 4 \
|
| 44 |
+
--d-ff 1024 \
|
| 45 |
+
--epochs 50 \
|
| 46 |
+
--batch-size 128 \
|
| 47 |
+
--lr 3e-4 \
|
| 48 |
+
--warmup-steps 500 \
|
| 49 |
+
--seed "$SEED" \
|
| 50 |
+
--device cuda
|
| 51 |
+
|
| 52 |
+
echo ""
|
| 53 |
+
echo "✅ Training complete for seed $SEED"
|
| 54 |
+
echo "Best checkpoint: $OUT_DIR/best.pt"
|
scripts/slurm/train_hybrid_direct.sbatch
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=hybrid_direct
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/hybrid_direct_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/hybrid_direct_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# DoVLA-Hybrid: DIRECT Scoring (NOT Pairwise)
|
| 16 |
+
# Expected: 45-48% baseline (vs 37% pairwise)
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/cvpr_hybrid_direct_model"
|
| 25 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 26 |
+
|
| 27 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 28 |
+
|
| 29 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 30 |
+
echo "DoVLA-Hybrid: DIRECT Scoring (FIXED!)"
|
| 31 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 32 |
+
echo ""
|
| 33 |
+
echo "KEY IMPROVEMENT:"
|
| 34 |
+
echo " OLD: Pairwise ranking → aggregate → 37%"
|
| 35 |
+
echo " NEW: Direct scoring → 45-48%"
|
| 36 |
+
echo ""
|
| 37 |
+
echo "Approach:"
|
| 38 |
+
echo " - Predict reward(action) directly"
|
| 39 |
+
echo " - Predict success(action) directly"
|
| 40 |
+
echo " - Select: argmax(success_prob * reward)"
|
| 41 |
+
echo ""
|
| 42 |
+
echo "Expected: 45-48% WITHOUT language"
|
| 43 |
+
echo "Then +language: 55-60% final"
|
| 44 |
+
echo ""
|
| 45 |
+
echo "Seed: $SEED"
|
| 46 |
+
echo ""
|
| 47 |
+
|
| 48 |
+
python scripts/train_hybrid_direct.py \
|
| 49 |
+
--dataset "$DATASET" \
|
| 50 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 51 |
+
--d-model 256 \
|
| 52 |
+
--n-heads 8 \
|
| 53 |
+
--n-layers 3 \
|
| 54 |
+
--d-ff 1024 \
|
| 55 |
+
--epochs 50 \
|
| 56 |
+
--batch-size 16 \
|
| 57 |
+
--lr 0.001 \
|
| 58 |
+
--weight-decay 0.01 \
|
| 59 |
+
--warmup-steps 500 \
|
| 60 |
+
--seed $SEED \
|
| 61 |
+
--device auto
|
| 62 |
+
|
| 63 |
+
echo ""
|
| 64 |
+
echo "✅ Hybrid training complete (seed $SEED)"
|
| 65 |
+
echo ""
|
scripts/slurm/train_maniskill_baseline_array.sbatch
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_baseline
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=28G
|
| 9 |
+
#SBATCH --time=02:00:00
|
| 10 |
+
#SBATCH --array=0-2%3
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
BASELINE="${BASELINE:?Set BASELINE}"
|
| 18 |
+
DATASET="${DATASET:?Set DATASET}"
|
| 19 |
+
RUN_ROOT="${RUN_ROOT:?Set RUN_ROOT}"
|
| 20 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 21 |
+
SEED="${SLURM_ARRAY_TASK_ID:-0}"
|
| 22 |
+
EPOCHS="${EPOCHS:-50}"
|
| 23 |
+
RECORDS_PER_GROUP="${RECORDS_PER_GROUP:-16}"
|
| 24 |
+
PAIR_SCOPE="same_state"
|
| 25 |
+
LOSS_ARGS=()
|
| 26 |
+
|
| 27 |
+
case "$BASELINE" in
|
| 28 |
+
cross_state_negatives)
|
| 29 |
+
PAIR_SCOPE="cross_state"
|
| 30 |
+
;;
|
| 31 |
+
random_negatives)
|
| 32 |
+
;;
|
| 33 |
+
world_model_auxiliary|no_rank_regret)
|
| 34 |
+
LOSS_ARGS+=(--loss-weight rank=0 --loss-weight regret=0)
|
| 35 |
+
;;
|
| 36 |
+
no_effect_head)
|
| 37 |
+
LOSS_ARGS+=(--loss-weight effect=0)
|
| 38 |
+
;;
|
| 39 |
+
label_only_counterfactual)
|
| 40 |
+
LOSS_ARGS+=(--loss-weight effect=0)
|
| 41 |
+
;;
|
| 42 |
+
expert_only_bc)
|
| 43 |
+
RECORDS_PER_GROUP=1
|
| 44 |
+
LOSS_ARGS+=(
|
| 45 |
+
--loss-weight effect=0
|
| 46 |
+
--loss-weight progress=0
|
| 47 |
+
--loss-weight rank=0
|
| 48 |
+
--loss-weight regret=0
|
| 49 |
+
)
|
| 50 |
+
;;
|
| 51 |
+
*)
|
| 52 |
+
echo "unsupported baseline: $BASELINE" >&2
|
| 53 |
+
exit 2
|
| 54 |
+
;;
|
| 55 |
+
esac
|
| 56 |
+
|
| 57 |
+
OUT_DIR="$RUN_ROOT/$BASELINE/seed_$SEED"
|
| 58 |
+
cd "$PROJECT_DIR"
|
| 59 |
+
mkdir -p "$OUT_DIR"
|
| 60 |
+
export OMP_NUM_THREADS=1
|
| 61 |
+
export OPENBLAS_NUM_THREADS=1
|
| 62 |
+
export MKL_NUM_THREADS=1
|
| 63 |
+
export DOVLA_TORCH_THREADS=1
|
| 64 |
+
|
| 65 |
+
"$PYTHON" scripts/train_dovla.py \
|
| 66 |
+
--dataset "$DATASET" \
|
| 67 |
+
--out "$OUT_DIR" \
|
| 68 |
+
--epochs "$EPOCHS" \
|
| 69 |
+
--batch-groups 32 \
|
| 70 |
+
--records-per-group "$RECORDS_PER_GROUP" \
|
| 71 |
+
--pair-count-per-group 32 \
|
| 72 |
+
--hidden-dim 256 \
|
| 73 |
+
--obs-dim 96 \
|
| 74 |
+
--lang-dim 64 \
|
| 75 |
+
--action-dim 8 \
|
| 76 |
+
--action-horizon 4 \
|
| 77 |
+
--effect-dim 32 \
|
| 78 |
+
--lr 0.001 \
|
| 79 |
+
--device cuda \
|
| 80 |
+
--seed "$SEED" \
|
| 81 |
+
--val-fraction 0.2 \
|
| 82 |
+
--objective legacy \
|
| 83 |
+
--pair-scope "$PAIR_SCOPE" \
|
| 84 |
+
"${LOSS_ARGS[@]}"
|
scripts/slurm/train_maniskill_collection_array.sbatch
ADDED
|
@@ -0,0 +1,114 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_multi_train
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=28G
|
| 9 |
+
#SBATCH --time=02:00:00
|
| 10 |
+
#SBATCH --array=0-5%6
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
DATASET="${DATASET:?Set DATASET to a CIL collection}"
|
| 18 |
+
RUN_ROOT="${RUN_ROOT:?Set RUN_ROOT}"
|
| 19 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 20 |
+
BACKBONE="${BACKBONE:-native}"
|
| 21 |
+
OBSERVATION_MODE="${OBSERVATION_MODE:-state}"
|
| 22 |
+
BACKBONE_MODEL="${BACKBONE_MODEL:-}"
|
| 23 |
+
BACKBONE_FEATURE_CACHE="${BACKBONE_FEATURE_CACHE:-}"
|
| 24 |
+
BACKBONE_FEATURE_BATCH_SIZE="${BACKBONE_FEATURE_BATCH_SIZE:-64}"
|
| 25 |
+
TASK_INDEX="${SLURM_ARRAY_TASK_ID:-0}"
|
| 26 |
+
OBJECTIVE_MODE="${OBJECTIVE_MODE:-paired}"
|
| 27 |
+
EPOCHS="${EPOCHS:-50}"
|
| 28 |
+
BATCH_GROUPS="${BATCH_GROUPS:-32}"
|
| 29 |
+
HIDDEN_DIM="${HIDDEN_DIM:-256}"
|
| 30 |
+
if [[ "$OBJECTIVE_MODE" == "field_only" ]]; then
|
| 31 |
+
SEED="$TASK_INDEX"
|
| 32 |
+
OBJECTIVE="${OBJECTIVE:-lattice_field}"
|
| 33 |
+
elif [[ "$OBJECTIVE_MODE" == "paired" ]]; then
|
| 34 |
+
SEED="$((TASK_INDEX / 2))"
|
| 35 |
+
if (( TASK_INDEX % 2 == 0 )); then
|
| 36 |
+
OBJECTIVE="lattice_field"
|
| 37 |
+
else
|
| 38 |
+
OBJECTIVE="legacy"
|
| 39 |
+
fi
|
| 40 |
+
else
|
| 41 |
+
echo "OBJECTIVE_MODE must be paired or field_only" >&2
|
| 42 |
+
exit 2
|
| 43 |
+
fi
|
| 44 |
+
OUT_DIR="$RUN_ROOT/$OBJECTIVE/seed_$SEED"
|
| 45 |
+
|
| 46 |
+
cd "$PROJECT_DIR"
|
| 47 |
+
mkdir -p "$OUT_DIR"
|
| 48 |
+
export OMP_NUM_THREADS=1
|
| 49 |
+
export OPENBLAS_NUM_THREADS=1
|
| 50 |
+
export MKL_NUM_THREADS=1
|
| 51 |
+
export DOVLA_TORCH_THREADS=1
|
| 52 |
+
|
| 53 |
+
if [[ "$BACKBONE" == "clip" ]]; then
|
| 54 |
+
[[ "$OBSERVATION_MODE" == "rgb" ]] || { echo "CLIP requires OBSERVATION_MODE=rgb" >&2; exit 2; }
|
| 55 |
+
[[ -n "$BACKBONE_MODEL" ]] || { echo "Set BACKBONE_MODEL for CLIP" >&2; exit 2; }
|
| 56 |
+
[[ -n "$BACKBONE_FEATURE_CACHE" ]] || { echo "Set BACKBONE_FEATURE_CACHE for CLIP" >&2; exit 2; }
|
| 57 |
+
fi
|
| 58 |
+
if [[ "$BACKBONE" == "clip" || "$OBSERVATION_MODE" == "rgb" ]]; then
|
| 59 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 60 |
+
SIF="${SIF:-$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif}"
|
| 61 |
+
CONTAINER_PYTHON="${CONTAINER_PYTHON:-$SCRATCH_ROOT/envs/maniskill/bin/python}"
|
| 62 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 63 |
+
PYTHON_COMMAND=(
|
| 64 |
+
apptainer exec --nv
|
| 65 |
+
--env "OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1,DOVLA_TORCH_THREADS=1,TRANSFORMERS_OFFLINE=1,HF_HUB_OFFLINE=1"
|
| 66 |
+
-B "$PROJECT_DIR:$PROJECT_DIR"
|
| 67 |
+
-B "/scratch/$USER:/scratch/$USER"
|
| 68 |
+
"$SIF" "$CONTAINER_PYTHON"
|
| 69 |
+
)
|
| 70 |
+
else
|
| 71 |
+
PYTHON_COMMAND=("$PYTHON")
|
| 72 |
+
fi
|
| 73 |
+
|
| 74 |
+
"${PYTHON_COMMAND[@]}" - <<PY
|
| 75 |
+
from dovla_cil.data.datasets import CILDataset
|
| 76 |
+
import torch
|
| 77 |
+
|
| 78 |
+
dataset = CILDataset("$DATASET")
|
| 79 |
+
assert len(dataset.group_ids) == int("${EXPECTED_GROUPS:-3500}"), len(dataset.group_ids)
|
| 80 |
+
assert len(dataset) == int("${EXPECTED_RECORDS:-56000}"), len(dataset)
|
| 81 |
+
assert torch.cuda.is_available()
|
| 82 |
+
print("objective=$OBJECTIVE seed=$SEED groups=", len(dataset.group_ids), "records=", len(dataset), torch.cuda.get_device_name(0))
|
| 83 |
+
PY
|
| 84 |
+
|
| 85 |
+
BACKBONE_ARGS=(--backbone "$BACKBONE")
|
| 86 |
+
if [[ "$BACKBONE" == "clip" ]]; then
|
| 87 |
+
BACKBONE_ARGS+=(
|
| 88 |
+
--backbone-model "$BACKBONE_MODEL"
|
| 89 |
+
--backbone-feature-cache "$BACKBONE_FEATURE_CACHE"
|
| 90 |
+
--backbone-feature-batch-size "$BACKBONE_FEATURE_BATCH_SIZE"
|
| 91 |
+
)
|
| 92 |
+
fi
|
| 93 |
+
|
| 94 |
+
"${PYTHON_COMMAND[@]}" scripts/train_dovla.py \
|
| 95 |
+
--dataset "$DATASET" \
|
| 96 |
+
--out "$OUT_DIR" \
|
| 97 |
+
--epochs "$EPOCHS" \
|
| 98 |
+
--batch-groups "$BATCH_GROUPS" \
|
| 99 |
+
--records-per-group 16 \
|
| 100 |
+
--pair-count-per-group 32 \
|
| 101 |
+
--hidden-dim "$HIDDEN_DIM" \
|
| 102 |
+
--obs-dim 96 \
|
| 103 |
+
--observation-mode "$OBSERVATION_MODE" \
|
| 104 |
+
--lang-dim 64 \
|
| 105 |
+
--action-dim 8 \
|
| 106 |
+
--action-horizon 4 \
|
| 107 |
+
--effect-dim 32 \
|
| 108 |
+
--lr 0.001 \
|
| 109 |
+
--device cuda \
|
| 110 |
+
--seed "$SEED" \
|
| 111 |
+
--val-fraction 0.2 \
|
| 112 |
+
--objective "$OBJECTIVE" \
|
| 113 |
+
--lattice-neighbors 32 \
|
| 114 |
+
"${BACKBONE_ARGS[@]}"
|
scripts/slurm/train_maniskill_collection_cpu_array.sbatch
ADDED
|
@@ -0,0 +1,61 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_cpu_train
|
| 3 |
+
#SBATCH --account=def-yalda_cpu
|
| 4 |
+
#SBATCH --partition=cpubase_bycore_b2
|
| 5 |
+
#SBATCH --nodes=1
|
| 6 |
+
#SBATCH --ntasks=1
|
| 7 |
+
#SBATCH --cpus-per-task=8
|
| 8 |
+
#SBATCH --mem=32G
|
| 9 |
+
#SBATCH --time=03:00:00
|
| 10 |
+
#SBATCH --array=0-2%3
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
DATASET="${DATASET:?Set DATASET to a CIL collection}"
|
| 18 |
+
RUN_ROOT="${RUN_ROOT:?Set RUN_ROOT}"
|
| 19 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 20 |
+
OBJECTIVE="${OBJECTIVE:-lattice_field}"
|
| 21 |
+
EPOCHS="${EPOCHS:-50}"
|
| 22 |
+
HIDDEN_DIM="${HIDDEN_DIM:-256}"
|
| 23 |
+
TASK_INDEX="${SLURM_ARRAY_TASK_ID:-0}"
|
| 24 |
+
SEED="${SEED_OVERRIDE:-$TASK_INDEX}"
|
| 25 |
+
OUT_DIR="$RUN_ROOT/$OBJECTIVE/seed_$SEED"
|
| 26 |
+
|
| 27 |
+
cd "$PROJECT_DIR"
|
| 28 |
+
mkdir -p outputs/hpc/logs "$OUT_DIR"
|
| 29 |
+
export OMP_NUM_THREADS=1
|
| 30 |
+
export OPENBLAS_NUM_THREADS=1
|
| 31 |
+
export MKL_NUM_THREADS=1
|
| 32 |
+
export DOVLA_TORCH_THREADS=1
|
| 33 |
+
|
| 34 |
+
"$PYTHON" - <<PY
|
| 35 |
+
from dovla_cil.data.datasets import CILDataset
|
| 36 |
+
|
| 37 |
+
dataset = CILDataset("$DATASET")
|
| 38 |
+
assert len(dataset.group_ids) == int("${EXPECTED_GROUPS:-3500}"), len(dataset.group_ids)
|
| 39 |
+
assert len(dataset) == int("${EXPECTED_RECORDS:-56000}"), len(dataset)
|
| 40 |
+
print("objective=$OBJECTIVE seed=$SEED device=cpu groups=", len(dataset.group_ids), "records=", len(dataset))
|
| 41 |
+
PY
|
| 42 |
+
|
| 43 |
+
"$PYTHON" scripts/train_dovla.py \
|
| 44 |
+
--dataset "$DATASET" \
|
| 45 |
+
--out "$OUT_DIR" \
|
| 46 |
+
--epochs "$EPOCHS" \
|
| 47 |
+
--batch-groups 32 \
|
| 48 |
+
--records-per-group 16 \
|
| 49 |
+
--pair-count-per-group 32 \
|
| 50 |
+
--hidden-dim "$HIDDEN_DIM" \
|
| 51 |
+
--obs-dim 96 \
|
| 52 |
+
--lang-dim 64 \
|
| 53 |
+
--action-dim 8 \
|
| 54 |
+
--action-horizon 4 \
|
| 55 |
+
--effect-dim 32 \
|
| 56 |
+
--lr 0.001 \
|
| 57 |
+
--device cpu \
|
| 58 |
+
--seed "$SEED" \
|
| 59 |
+
--val-fraction 0.2 \
|
| 60 |
+
--objective "$OBJECTIVE" \
|
| 61 |
+
--lattice-neighbors 32
|
scripts/slurm/train_maniskill_debug.sbatch
ADDED
|
@@ -0,0 +1,86 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_train_debug
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=00:20:00
|
| 10 |
+
#SBATCH --output=outputs/hpc/logs/%x_%j.out
|
| 11 |
+
#SBATCH --error=outputs/hpc/logs/%x_%j.err
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 16 |
+
DATASET="${DATASET:-$PROJECT_DIR/outputs/hpc/maniskill_debug_cil}"
|
| 17 |
+
RUN_ROOT="${RUN_ROOT:-$PROJECT_DIR/outputs/hpc/maniskill_debug_runs}"
|
| 18 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 19 |
+
export DATASET RUN_ROOT
|
| 20 |
+
|
| 21 |
+
cd "$PROJECT_DIR"
|
| 22 |
+
mkdir -p outputs/hpc/logs "$RUN_ROOT"
|
| 23 |
+
|
| 24 |
+
export OMP_NUM_THREADS=1
|
| 25 |
+
export OPENBLAS_NUM_THREADS=1
|
| 26 |
+
export MKL_NUM_THREADS=1
|
| 27 |
+
export DOVLA_TORCH_THREADS=1
|
| 28 |
+
|
| 29 |
+
test -f "$DATASET/manifest.json"
|
| 30 |
+
"$PYTHON" - <<'PY'
|
| 31 |
+
import json
|
| 32 |
+
import os
|
| 33 |
+
from pathlib import Path
|
| 34 |
+
|
| 35 |
+
import torch
|
| 36 |
+
|
| 37 |
+
manifest = json.loads((Path(os.environ["DATASET"]) / "manifest.json").read_text())
|
| 38 |
+
print("torch", torch.__version__, "cuda", torch.cuda.is_available())
|
| 39 |
+
if torch.cuda.is_available():
|
| 40 |
+
print("gpu", torch.cuda.get_device_name(0))
|
| 41 |
+
print("dataset_records", manifest["record_count"], "groups", manifest["group_count"])
|
| 42 |
+
assert manifest["record_count"] > 0 and manifest["group_count"] >= 4
|
| 43 |
+
PY
|
| 44 |
+
|
| 45 |
+
COMMON_ARGS=(
|
| 46 |
+
--dataset "$DATASET"
|
| 47 |
+
--epochs 10
|
| 48 |
+
--batch-groups 2
|
| 49 |
+
--records-per-group 4
|
| 50 |
+
--pair-count-per-group 6
|
| 51 |
+
--hidden-dim 128
|
| 52 |
+
--obs-dim 96
|
| 53 |
+
--lang-dim 64
|
| 54 |
+
--action-dim 7
|
| 55 |
+
--action-horizon 4
|
| 56 |
+
--effect-dim 16
|
| 57 |
+
--lr 0.001
|
| 58 |
+
--device cuda
|
| 59 |
+
--seed 0
|
| 60 |
+
--val-fraction 0.25
|
| 61 |
+
)
|
| 62 |
+
|
| 63 |
+
"$PYTHON" scripts/train_dovla.py \
|
| 64 |
+
"${COMMON_ARGS[@]}" \
|
| 65 |
+
--objective lattice_field \
|
| 66 |
+
--lattice-neighbors 2 \
|
| 67 |
+
--out "$RUN_ROOT/lattice_field"
|
| 68 |
+
|
| 69 |
+
"$PYTHON" scripts/train_dovla.py \
|
| 70 |
+
"${COMMON_ARGS[@]}" \
|
| 71 |
+
--objective legacy \
|
| 72 |
+
--out "$RUN_ROOT/legacy"
|
| 73 |
+
|
| 74 |
+
"$PYTHON" - <<'PY'
|
| 75 |
+
import json
|
| 76 |
+
import os
|
| 77 |
+
from pathlib import Path
|
| 78 |
+
|
| 79 |
+
root = Path(os.environ["RUN_ROOT"])
|
| 80 |
+
summary = {}
|
| 81 |
+
for name in ("lattice_field", "legacy"):
|
| 82 |
+
metrics = json.loads((root / name / "metrics.json").read_text())
|
| 83 |
+
summary[name] = metrics["best"]
|
| 84 |
+
(root / "comparison.json").write_text(json.dumps(summary, indent=2, sort_keys=True) + "\n")
|
| 85 |
+
print(json.dumps(summary, indent=2, sort_keys=True))
|
| 86 |
+
PY
|
scripts/slurm/train_maniskill_full_array.sbatch
ADDED
|
@@ -0,0 +1,74 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_train
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=01:00:00
|
| 10 |
+
#SBATCH --array=0-5%6
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
DATASET="${DATASET:-$PROJECT_DIR/outputs/hpc/maniskill_full_k16_n1000_seed0}"
|
| 18 |
+
RUN_ROOT="${RUN_ROOT:-$PROJECT_DIR/outputs/hpc/maniskill_full_runs}"
|
| 19 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 20 |
+
|
| 21 |
+
TASK_INDEX="${SLURM_ARRAY_TASK_ID:-0}"
|
| 22 |
+
SEED="$((TASK_INDEX / 2))"
|
| 23 |
+
if (( TASK_INDEX % 2 == 0 )); then
|
| 24 |
+
OBJECTIVE="lattice_field"
|
| 25 |
+
else
|
| 26 |
+
OBJECTIVE="legacy"
|
| 27 |
+
fi
|
| 28 |
+
OUT_DIR="$RUN_ROOT/$OBJECTIVE/seed_$SEED"
|
| 29 |
+
|
| 30 |
+
cd "$PROJECT_DIR"
|
| 31 |
+
mkdir -p outputs/hpc/logs "$OUT_DIR"
|
| 32 |
+
|
| 33 |
+
export OMP_NUM_THREADS=1
|
| 34 |
+
export OPENBLAS_NUM_THREADS=1
|
| 35 |
+
export MKL_NUM_THREADS=1
|
| 36 |
+
export DOVLA_TORCH_THREADS=1
|
| 37 |
+
|
| 38 |
+
test -f "$DATASET/manifest.json"
|
| 39 |
+
"$PYTHON" - <<PY
|
| 40 |
+
import json
|
| 41 |
+
from pathlib import Path
|
| 42 |
+
import torch
|
| 43 |
+
|
| 44 |
+
manifest = json.loads((Path("$DATASET") / "manifest.json").read_text())
|
| 45 |
+
assert manifest["group_count"] == 1000
|
| 46 |
+
assert manifest["record_count"] == 16000
|
| 47 |
+
assert torch.cuda.is_available()
|
| 48 |
+
print(
|
| 49 |
+
"objective=$OBJECTIVE seed=$SEED",
|
| 50 |
+
"gpu=", torch.cuda.get_device_name(0),
|
| 51 |
+
"groups=", manifest["group_count"],
|
| 52 |
+
"records=", manifest["record_count"],
|
| 53 |
+
)
|
| 54 |
+
PY
|
| 55 |
+
|
| 56 |
+
"$PYTHON" scripts/train_dovla.py \
|
| 57 |
+
--dataset "$DATASET" \
|
| 58 |
+
--out "$OUT_DIR" \
|
| 59 |
+
--epochs 50 \
|
| 60 |
+
--batch-groups 32 \
|
| 61 |
+
--records-per-group 16 \
|
| 62 |
+
--pair-count-per-group 32 \
|
| 63 |
+
--hidden-dim 256 \
|
| 64 |
+
--obs-dim 96 \
|
| 65 |
+
--lang-dim 64 \
|
| 66 |
+
--action-dim 7 \
|
| 67 |
+
--action-horizon 4 \
|
| 68 |
+
--effect-dim 32 \
|
| 69 |
+
--lr 0.001 \
|
| 70 |
+
--device cuda \
|
| 71 |
+
--seed "$SEED" \
|
| 72 |
+
--val-fraction 0.2 \
|
| 73 |
+
--objective "$OBJECTIVE" \
|
| 74 |
+
--lattice-neighbors 32
|
scripts/slurm/train_maniskill_scaling_array.sbatch
ADDED
|
@@ -0,0 +1,82 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_scale_train
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=4
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=24G
|
| 9 |
+
#SBATCH --time=01:00:00
|
| 10 |
+
#SBATCH --array=0-2%3
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
K="${K:?export K before submission}"
|
| 18 |
+
NUM_GROUPS="${NUM_GROUPS:?export NUM_GROUPS before submission}"
|
| 19 |
+
TOTAL_RECORDS="${TOTAL_RECORDS:-16000}"
|
| 20 |
+
DATASET="${DATASET:?export DATASET before submission}"
|
| 21 |
+
RUN_ROOT="${RUN_ROOT:-$PROJECT_DIR/outputs/hpc/maniskill_scaling_runs}"
|
| 22 |
+
PYTHON="${PYTHON:-$PROJECT_DIR/.venv/bin/python}"
|
| 23 |
+
SEED="${SLURM_ARRAY_TASK_ID:-0}"
|
| 24 |
+
BATCH_GROUPS="$((512 / K))"
|
| 25 |
+
PAIR_COUNT="$((2 * K))"
|
| 26 |
+
OUT_DIR="$RUN_ROOT/k_$K/seed_$SEED"
|
| 27 |
+
|
| 28 |
+
if (( K <= 0 || 512 % K != 0 )); then
|
| 29 |
+
echo "K must be a positive divisor of 512" >&2
|
| 30 |
+
exit 2
|
| 31 |
+
fi
|
| 32 |
+
if (( NUM_GROUPS * K != TOTAL_RECORDS )); then
|
| 33 |
+
echo "fixed-budget invariant failed: NUM_GROUPS*K != TOTAL_RECORDS" >&2
|
| 34 |
+
exit 2
|
| 35 |
+
fi
|
| 36 |
+
|
| 37 |
+
cd "$PROJECT_DIR"
|
| 38 |
+
mkdir -p outputs/hpc/logs "$OUT_DIR"
|
| 39 |
+
|
| 40 |
+
export OMP_NUM_THREADS=1
|
| 41 |
+
export OPENBLAS_NUM_THREADS=1
|
| 42 |
+
export MKL_NUM_THREADS=1
|
| 43 |
+
export DOVLA_TORCH_THREADS=1
|
| 44 |
+
|
| 45 |
+
test -f "$DATASET/manifest.json"
|
| 46 |
+
"$PYTHON" - <<PY
|
| 47 |
+
import json
|
| 48 |
+
from pathlib import Path
|
| 49 |
+
import torch
|
| 50 |
+
|
| 51 |
+
manifest = json.loads((Path("$DATASET") / "manifest.json").read_text())
|
| 52 |
+
assert manifest["group_count"] == int("$NUM_GROUPS")
|
| 53 |
+
assert manifest["record_count"] == int("$TOTAL_RECORDS")
|
| 54 |
+
assert manifest["k"] == int("$K")
|
| 55 |
+
assert torch.cuda.is_available()
|
| 56 |
+
print(
|
| 57 |
+
"K=$K seed=$SEED batch_groups=$BATCH_GROUPS",
|
| 58 |
+
"gpu=", torch.cuda.get_device_name(0),
|
| 59 |
+
"groups=", manifest["group_count"],
|
| 60 |
+
"records=", manifest["record_count"],
|
| 61 |
+
)
|
| 62 |
+
PY
|
| 63 |
+
|
| 64 |
+
"$PYTHON" scripts/train_dovla.py \
|
| 65 |
+
--dataset "$DATASET" \
|
| 66 |
+
--out "$OUT_DIR" \
|
| 67 |
+
--epochs 50 \
|
| 68 |
+
--batch-groups "$BATCH_GROUPS" \
|
| 69 |
+
--records-per-group "$K" \
|
| 70 |
+
--pair-count-per-group "$PAIR_COUNT" \
|
| 71 |
+
--hidden-dim 256 \
|
| 72 |
+
--obs-dim 96 \
|
| 73 |
+
--lang-dim 64 \
|
| 74 |
+
--action-dim 7 \
|
| 75 |
+
--action-horizon 4 \
|
| 76 |
+
--effect-dim 32 \
|
| 77 |
+
--lr 0.001 \
|
| 78 |
+
--device cuda \
|
| 79 |
+
--seed "$SEED" \
|
| 80 |
+
--val-fraction 0.2 \
|
| 81 |
+
--objective lattice_field \
|
| 82 |
+
--lattice-neighbors 32
|
scripts/slurm/train_maniskill_visual_array.sbatch
ADDED
|
@@ -0,0 +1,79 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_ms_rgb_train
|
| 3 |
+
#SBATCH --account=def-yalda_gpu
|
| 4 |
+
#SBATCH --nodes=1
|
| 5 |
+
#SBATCH --ntasks=1
|
| 6 |
+
#SBATCH --cpus-per-task=8
|
| 7 |
+
#SBATCH --gres=gpu:nvidia_h100_80gb_hbm3_1g.10gb:1
|
| 8 |
+
#SBATCH --mem=28G
|
| 9 |
+
#SBATCH --time=03:00:00
|
| 10 |
+
#SBATCH --array=0-2%3
|
| 11 |
+
#SBATCH --output=outputs/hpc/logs/%x_%A_%a.out
|
| 12 |
+
#SBATCH --error=outputs/hpc/logs/%x_%A_%a.err
|
| 13 |
+
|
| 14 |
+
set -euo pipefail
|
| 15 |
+
|
| 16 |
+
PROJECT_DIR="${PROJECT_DIR:-$SLURM_SUBMIT_DIR}"
|
| 17 |
+
DATASET="${DATASET:?Set DATASET to a rendered CIL collection}"
|
| 18 |
+
RUN_ROOT="${RUN_ROOT:?Set RUN_ROOT}"
|
| 19 |
+
SCRATCH_ROOT="/scratch/$USER/dovla"
|
| 20 |
+
SIF="${SIF:-$SCRATCH_ROOT/containers/pytorch_2.7.1_cuda12.8.sif}"
|
| 21 |
+
PYTHON="${PYTHON:-$SCRATCH_ROOT/envs/maniskill/bin/python}"
|
| 22 |
+
SEED="${SLURM_ARRAY_TASK_ID:-0}"
|
| 23 |
+
EPOCHS="${EPOCHS:-10}"
|
| 24 |
+
BATCH_GROUPS="${BATCH_GROUPS:-8}"
|
| 25 |
+
HIDDEN_DIM="${HIDDEN_DIM:-256}"
|
| 26 |
+
OUT_DIR="$RUN_ROOT/lattice_field/seed_$SEED"
|
| 27 |
+
|
| 28 |
+
cd "$PROJECT_DIR"
|
| 29 |
+
mkdir -p "$OUT_DIR"
|
| 30 |
+
module load StdEnv/2023 apptainer/1.4.5
|
| 31 |
+
export OMP_NUM_THREADS=1
|
| 32 |
+
export OPENBLAS_NUM_THREADS=1
|
| 33 |
+
export MKL_NUM_THREADS=1
|
| 34 |
+
export DOVLA_TORCH_THREADS=1
|
| 35 |
+
RUNTIME=(
|
| 36 |
+
apptainer exec --nv
|
| 37 |
+
--env "OMP_NUM_THREADS=1,OPENBLAS_NUM_THREADS=1,MKL_NUM_THREADS=1,DOVLA_TORCH_THREADS=1"
|
| 38 |
+
-B "$PROJECT_DIR:$PROJECT_DIR"
|
| 39 |
+
-B "/scratch/$USER:/scratch/$USER"
|
| 40 |
+
"$SIF"
|
| 41 |
+
"$PYTHON"
|
| 42 |
+
)
|
| 43 |
+
|
| 44 |
+
"${RUNTIME[@]}" - <<PY
|
| 45 |
+
from dovla_cil.data.datasets import CILDataset
|
| 46 |
+
import torch
|
| 47 |
+
|
| 48 |
+
dataset = CILDataset("$DATASET")
|
| 49 |
+
assert dataset.group_ids and len(dataset) > 0
|
| 50 |
+
assert all(record.observation_ref for record in dataset.records[: min(256, len(dataset))])
|
| 51 |
+
assert torch.cuda.is_available()
|
| 52 |
+
print(
|
| 53 |
+
"visual seed=$SEED",
|
| 54 |
+
"gpu=", torch.cuda.get_device_name(0),
|
| 55 |
+
"groups=", len(dataset.group_ids),
|
| 56 |
+
"records=", len(dataset),
|
| 57 |
+
)
|
| 58 |
+
PY
|
| 59 |
+
|
| 60 |
+
"${RUNTIME[@]}" scripts/train_dovla.py \
|
| 61 |
+
--dataset "$DATASET" \
|
| 62 |
+
--out "$OUT_DIR" \
|
| 63 |
+
--epochs "$EPOCHS" \
|
| 64 |
+
--batch-groups "$BATCH_GROUPS" \
|
| 65 |
+
--records-per-group 16 \
|
| 66 |
+
--pair-count-per-group 32 \
|
| 67 |
+
--hidden-dim "$HIDDEN_DIM" \
|
| 68 |
+
--obs-dim 96 \
|
| 69 |
+
--observation-mode rgb \
|
| 70 |
+
--lang-dim 64 \
|
| 71 |
+
--action-dim 8 \
|
| 72 |
+
--action-horizon 4 \
|
| 73 |
+
--effect-dim 32 \
|
| 74 |
+
--lr 0.001 \
|
| 75 |
+
--device cuda \
|
| 76 |
+
--seed "$SEED" \
|
| 77 |
+
--val-fraction 0.2 \
|
| 78 |
+
--objective lattice_field \
|
| 79 |
+
--lattice-neighbors 32
|
scripts/slurm/train_transformer.sbatch
ADDED
|
@@ -0,0 +1,72 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=dovla_transformer
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/transformer_train_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/transformer_train_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# DoVLA-Transformer: Pure Transformer Architecture (BREAKTHROUGH)
|
| 16 |
+
# Expected: 42-47% success (vs 38.43% baseline, 36.31% failed Enhanced)
|
| 17 |
+
|
| 18 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 19 |
+
cd "$PROJECT_DIR"
|
| 20 |
+
|
| 21 |
+
source .venv/bin/activate
|
| 22 |
+
|
| 23 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 24 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/cvpr_transformer_model"
|
| 25 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 26 |
+
|
| 27 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 28 |
+
|
| 29 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 30 |
+
echo "DoVLA-Transformer: BREAKTHROUGH Architecture"
|
| 31 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 32 |
+
echo ""
|
| 33 |
+
echo "Pure Transformer Components:"
|
| 34 |
+
echo " - Multi-head self-attention (8 heads)"
|
| 35 |
+
echo " - Cross-attention for obs-lang fusion"
|
| 36 |
+
echo " - 3 Transformer encoder layers"
|
| 37 |
+
echo " - Positional encoding"
|
| 38 |
+
echo " - Residual connections everywhere"
|
| 39 |
+
echo ""
|
| 40 |
+
echo "Key Improvements:"
|
| 41 |
+
echo " - Higher LR: 0.001 (vs 0.0003 failed Enhanced)"
|
| 42 |
+
echo " - Warmup scheduler: 500 steps"
|
| 43 |
+
echo " - No custom GNN (proven Transformer only)"
|
| 44 |
+
echo " - Proper gradient flow (residuals)"
|
| 45 |
+
echo ""
|
| 46 |
+
echo "Dataset: 3,500 groups (fair comparison)"
|
| 47 |
+
echo "Seed: $SEED"
|
| 48 |
+
echo ""
|
| 49 |
+
echo "Expected: 42-47% success"
|
| 50 |
+
echo "vs Baseline: 38.43%"
|
| 51 |
+
echo "vs Enhanced (failed): 36.31%"
|
| 52 |
+
echo ""
|
| 53 |
+
|
| 54 |
+
python scripts/train_dovla_transformer.py \
|
| 55 |
+
--dataset "$DATASET" \
|
| 56 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 57 |
+
--d-model 256 \
|
| 58 |
+
--n-heads 8 \
|
| 59 |
+
--n-layers 3 \
|
| 60 |
+
--d-ff 1024 \
|
| 61 |
+
--epochs 50 \
|
| 62 |
+
--batch-size 16 \
|
| 63 |
+
--lr 0.001 \
|
| 64 |
+
--weight-decay 0.01 \
|
| 65 |
+
--warmup-steps 500 \
|
| 66 |
+
--seed $SEED \
|
| 67 |
+
--device auto
|
| 68 |
+
|
| 69 |
+
echo ""
|
| 70 |
+
echo "✅ Transformer training complete (seed $SEED)"
|
| 71 |
+
echo ""
|
| 72 |
+
echo "Next: Evaluate and compare"
|
scripts/slurm/train_transformer_lang.sbatch
ADDED
|
@@ -0,0 +1,68 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/bin/bash
|
| 2 |
+
#SBATCH --job-name=transformer_lang
|
| 3 |
+
#SBATCH --nodes=1
|
| 4 |
+
#SBATCH --ntasks=1
|
| 5 |
+
#SBATCH --cpus-per-task=8
|
| 6 |
+
#SBATCH --gres=gpu:1
|
| 7 |
+
#SBATCH --mem=64000M
|
| 8 |
+
#SBATCH --time=48:00:00
|
| 9 |
+
#SBATCH --output=logs/transformer_lang_%A_%a.out
|
| 10 |
+
#SBATCH --error=logs/transformer_lang_%A_%a.err
|
| 11 |
+
#SBATCH --array=0-2
|
| 12 |
+
|
| 13 |
+
set -euo pipefail
|
| 14 |
+
|
| 15 |
+
# DoVLA-Transformer WITH LANGUAGE
|
| 16 |
+
# Expected: 50-55% (from 42-44% baseline)
|
| 17 |
+
# Improvement: +8-11%
|
| 18 |
+
|
| 19 |
+
PROJECT_DIR="${PROJECT_DIR:-$PWD}"
|
| 20 |
+
cd "$PROJECT_DIR"
|
| 21 |
+
|
| 22 |
+
source .venv/bin/activate
|
| 23 |
+
|
| 24 |
+
DATASET="/scratch/$USER/dovla/experiments/maniskill_presuccess_six_task_collection"
|
| 25 |
+
EMBEDDINGS="/scratch/$USER/dovla/experiments/instruction_embeddings.pkl"
|
| 26 |
+
OUT_DIR="/scratch/$USER/dovla/experiments/cvpr_transformer_lang_model"
|
| 27 |
+
SEED=$SLURM_ARRAY_TASK_ID
|
| 28 |
+
|
| 29 |
+
mkdir -p "$OUT_DIR/seed_$SEED" logs
|
| 30 |
+
|
| 31 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 32 |
+
echo "DoVLA-Transformer WITH LANGUAGE"
|
| 33 |
+
echo "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "=" "="
|
| 34 |
+
echo ""
|
| 35 |
+
echo "NEW FEATURE: Instruction embeddings (768-dim)"
|
| 36 |
+
echo " - Baseline (no language): 42-44%"
|
| 37 |
+
echo " - WITH language: 50-55% expected"
|
| 38 |
+
echo " - Improvement: +8-11%"
|
| 39 |
+
echo ""
|
| 40 |
+
echo "Architecture:"
|
| 41 |
+
echo " - Pure Transformer (8 heads, 3 layers)"
|
| 42 |
+
echo " - Language dimension: 768"
|
| 43 |
+
echo " - Cross-attention: obs + lang → context"
|
| 44 |
+
echo ""
|
| 45 |
+
echo "Dataset: 3,500 groups"
|
| 46 |
+
echo "Seed: $SEED"
|
| 47 |
+
echo ""
|
| 48 |
+
|
| 49 |
+
python scripts/train_transformer_with_language.py \
|
| 50 |
+
--dataset "$DATASET" \
|
| 51 |
+
--embeddings "$EMBEDDINGS" \
|
| 52 |
+
--out "$OUT_DIR/seed_$SEED" \
|
| 53 |
+
--d-model 256 \
|
| 54 |
+
--n-heads 8 \
|
| 55 |
+
--n-layers 3 \
|
| 56 |
+
--d-ff 1024 \
|
| 57 |
+
--epochs 50 \
|
| 58 |
+
--batch-size 16 \
|
| 59 |
+
--lr 0.001 \
|
| 60 |
+
--weight-decay 0.01 \
|
| 61 |
+
--warmup-steps 500 \
|
| 62 |
+
--seed $SEED \
|
| 63 |
+
--device auto
|
| 64 |
+
|
| 65 |
+
echo ""
|
| 66 |
+
echo "✅ Training with language complete (seed $SEED)"
|
| 67 |
+
echo ""
|
| 68 |
+
echo "Next: Evaluate and compare"
|
scripts/smoke_full_pipeline.py
ADDED
|
@@ -0,0 +1,155 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import argparse
|
| 5 |
+
import contextlib
|
| 6 |
+
import io
|
| 7 |
+
import os
|
| 8 |
+
import shutil
|
| 9 |
+
import sys
|
| 10 |
+
from pathlib import Path
|
| 11 |
+
|
| 12 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 13 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 14 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 15 |
+
|
| 16 |
+
from dovla_cil.eval.causalstress import ( # noqa: E402
|
| 17 |
+
CausalStressBenchmark,
|
| 18 |
+
CausalStressConfig,
|
| 19 |
+
write_metrics_json,
|
| 20 |
+
)
|
| 21 |
+
from dovla_cil.experiments.reports import generate_dataset_report, generate_eval_report
|
| 22 |
+
from dovla_cil.generation.pipeline import generate_cil_dataset, print_generation_summary
|
| 23 |
+
from dovla_cil.tasks.library import ToyTaskLibrary
|
| 24 |
+
from dovla_cil.training.trainer import DoVLATrainer, TrainerConfig
|
| 25 |
+
from dovla_cil.utils.io import ensure_dir, write_jsonl
|
| 26 |
+
from scripts.inspect_shard import main as inspect_main
|
| 27 |
+
|
| 28 |
+
|
| 29 |
+
def main(argv: list[str] | None = None) -> int:
|
| 30 |
+
parser = argparse.ArgumentParser(description="Run the full local DoVLA-CIL smoke pipeline.")
|
| 31 |
+
parser.add_argument("--out", type=Path, default=Path("outputs/smoke_full"))
|
| 32 |
+
parser.add_argument("--num-tasks", type=int, default=3)
|
| 33 |
+
parser.add_argument("--states-per-task", type=int, default=4)
|
| 34 |
+
parser.add_argument("--k", type=int, default=4)
|
| 35 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 36 |
+
parser.add_argument("--shard-size", type=int, default=32)
|
| 37 |
+
parser.add_argument("--epochs", type=int, default=1)
|
| 38 |
+
parser.add_argument("--batch-groups", type=int, default=2)
|
| 39 |
+
parser.add_argument("--records-per-group", type=int, default=4)
|
| 40 |
+
parser.add_argument("--hidden-dim", type=int, default=64)
|
| 41 |
+
parser.add_argument("--eval-num-tasks", type=int, default=6)
|
| 42 |
+
parser.add_argument("--device", default="cpu")
|
| 43 |
+
parser.add_argument("--no-clean", action="store_true", help="Do not remove an existing output dir.")
|
| 44 |
+
args = parser.parse_args(argv)
|
| 45 |
+
|
| 46 |
+
if args.num_tasks <= 0 or args.states_per_task <= 0 or args.k <= 0:
|
| 47 |
+
raise ValueError("num-tasks, states-per-task, and k must be positive")
|
| 48 |
+
|
| 49 |
+
os.environ.setdefault("OPENCLAUDE_MOCK", "1")
|
| 50 |
+
if args.out.exists() and not args.no_clean:
|
| 51 |
+
shutil.rmtree(args.out)
|
| 52 |
+
output_dir = ensure_dir(args.out)
|
| 53 |
+
dataset_dir = output_dir / "cil_toy"
|
| 54 |
+
train_dir = output_dir / "train"
|
| 55 |
+
report_dir = output_dir / "dataset_report"
|
| 56 |
+
eval_metrics_path = output_dir / "causalstress" / "metrics.json"
|
| 57 |
+
eval_report_dir = output_dir / "eval_report"
|
| 58 |
+
inspect_path = output_dir / "inspect.txt"
|
| 59 |
+
task_path = output_dir / "tasks.jsonl"
|
| 60 |
+
|
| 61 |
+
print("1. Loading built-in toy tasks")
|
| 62 |
+
tasks = ToyTaskLibrary().list(args.num_tasks)
|
| 63 |
+
write_jsonl((task.to_dict() for task in tasks), task_path)
|
| 64 |
+
print(f" tasks: {task_path}")
|
| 65 |
+
|
| 66 |
+
print("2. Generating CIL dataset")
|
| 67 |
+
generation_summary = generate_cil_dataset(
|
| 68 |
+
backend="toy",
|
| 69 |
+
tasks=tasks,
|
| 70 |
+
out_dir=dataset_dir,
|
| 71 |
+
num_states_per_task=args.states_per_task,
|
| 72 |
+
k=args.k,
|
| 73 |
+
seed=args.seed,
|
| 74 |
+
shard_size=args.shard_size,
|
| 75 |
+
inline_observations=True,
|
| 76 |
+
)
|
| 77 |
+
print_generation_summary(generation_summary)
|
| 78 |
+
|
| 79 |
+
print("3. Inspecting dataset")
|
| 80 |
+
inspect_output = _capture_stdout(
|
| 81 |
+
lambda: inspect_main([str(dataset_dir), "--max-rows", str(args.k)])
|
| 82 |
+
)
|
| 83 |
+
inspect_path.write_text(inspect_output, encoding="utf-8")
|
| 84 |
+
print(inspect_output.rstrip())
|
| 85 |
+
|
| 86 |
+
print("4. Training DoVLA for one smoke epoch")
|
| 87 |
+
trainer_result = DoVLATrainer(
|
| 88 |
+
TrainerConfig(
|
| 89 |
+
dataset_dir=dataset_dir,
|
| 90 |
+
output_dir=train_dir,
|
| 91 |
+
epochs=args.epochs,
|
| 92 |
+
batch_groups=args.batch_groups,
|
| 93 |
+
records_per_group=args.records_per_group,
|
| 94 |
+
pair_count_per_group=args.records_per_group,
|
| 95 |
+
hidden_dim=args.hidden_dim,
|
| 96 |
+
learning_rate=1e-3,
|
| 97 |
+
device=args.device,
|
| 98 |
+
seed=args.seed,
|
| 99 |
+
val_fraction=0.25,
|
| 100 |
+
)
|
| 101 |
+
).train()
|
| 102 |
+
print(f" checkpoints: {train_dir}")
|
| 103 |
+
print(f" best metrics: {trainer_result.get('best', {})}")
|
| 104 |
+
|
| 105 |
+
print("5. Evaluating CausalStress")
|
| 106 |
+
eval_config = CausalStressConfig(
|
| 107 |
+
backend="toy",
|
| 108 |
+
num_tasks=args.eval_num_tasks,
|
| 109 |
+
k=args.k,
|
| 110 |
+
seed=args.seed,
|
| 111 |
+
)
|
| 112 |
+
metrics = CausalStressBenchmark(eval_config).evaluate(train_dir / "best.pt", device=args.device)
|
| 113 |
+
metrics["config"] = {
|
| 114 |
+
"backend": "toy",
|
| 115 |
+
"checkpoint": str(train_dir / "best.pt"),
|
| 116 |
+
"num_tasks": args.eval_num_tasks,
|
| 117 |
+
"k": args.k,
|
| 118 |
+
"seed": args.seed,
|
| 119 |
+
}
|
| 120 |
+
write_metrics_json(metrics, eval_metrics_path)
|
| 121 |
+
print(f" metrics: {eval_metrics_path}")
|
| 122 |
+
|
| 123 |
+
print("6. Writing dataset report")
|
| 124 |
+
generate_dataset_report(dataset_dir, report_dir, sample_groups=3, seed=args.seed)
|
| 125 |
+
print(f" dataset report: {report_dir}")
|
| 126 |
+
|
| 127 |
+
print("7. Writing evaluation report")
|
| 128 |
+
generate_eval_report([eval_metrics_path], eval_report_dir, experiment_name="smoke_full")
|
| 129 |
+
print(f" eval report: {eval_report_dir}")
|
| 130 |
+
|
| 131 |
+
print("8. Final paths")
|
| 132 |
+
for label, path in {
|
| 133 |
+
"root": output_dir,
|
| 134 |
+
"tasks": task_path,
|
| 135 |
+
"dataset": dataset_dir,
|
| 136 |
+
"inspect": inspect_path,
|
| 137 |
+
"train": train_dir,
|
| 138 |
+
"checkpoint": train_dir / "best.pt",
|
| 139 |
+
"causalstress_metrics": eval_metrics_path,
|
| 140 |
+
"dataset_report": report_dir,
|
| 141 |
+
"eval_report": eval_report_dir,
|
| 142 |
+
}.items():
|
| 143 |
+
print(f" {label}: {path}")
|
| 144 |
+
return 0
|
| 145 |
+
|
| 146 |
+
|
| 147 |
+
def _capture_stdout(callback) -> str:
|
| 148 |
+
buffer = io.StringIO()
|
| 149 |
+
with contextlib.redirect_stdout(buffer):
|
| 150 |
+
callback()
|
| 151 |
+
return buffer.getvalue()
|
| 152 |
+
|
| 153 |
+
|
| 154 |
+
if __name__ == "__main__":
|
| 155 |
+
raise SystemExit(main())
|
scripts/smoke_smolvla_checkpoint.py
ADDED
|
@@ -0,0 +1,169 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python3
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import argparse
|
| 5 |
+
import importlib.metadata
|
| 6 |
+
import json
|
| 7 |
+
import sys
|
| 8 |
+
import time
|
| 9 |
+
from pathlib import Path
|
| 10 |
+
from typing import Any
|
| 11 |
+
|
| 12 |
+
ROOT = Path(__file__).resolve().parents[1]
|
| 13 |
+
if str(ROOT) not in sys.path:
|
| 14 |
+
sys.path.insert(0, str(ROOT))
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def build_parser() -> argparse.ArgumentParser:
|
| 18 |
+
parser = argparse.ArgumentParser(
|
| 19 |
+
description="Load a local SmolVLA checkpoint and write a reproducible smoke manifest."
|
| 20 |
+
)
|
| 21 |
+
parser.add_argument("--checkpoint", type=Path, required=True)
|
| 22 |
+
parser.add_argument(
|
| 23 |
+
"--vlm-metadata",
|
| 24 |
+
type=Path,
|
| 25 |
+
help=(
|
| 26 |
+
"Local SmolVLM config/tokenizer directory. Required for an offline weight-loading "
|
| 27 |
+
"smoke test."
|
| 28 |
+
),
|
| 29 |
+
)
|
| 30 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 31 |
+
parser.add_argument("--device", default="auto", choices=("auto", "cpu", "cuda"))
|
| 32 |
+
parser.add_argument(
|
| 33 |
+
"--metadata-only",
|
| 34 |
+
action="store_true",
|
| 35 |
+
help="Validate files and package availability without allocating model weights.",
|
| 36 |
+
)
|
| 37 |
+
return parser
|
| 38 |
+
|
| 39 |
+
|
| 40 |
+
def smoke_checkpoint(
|
| 41 |
+
checkpoint: Path,
|
| 42 |
+
*,
|
| 43 |
+
device: str = "auto",
|
| 44 |
+
metadata_only: bool = False,
|
| 45 |
+
vlm_metadata: Path | None = None,
|
| 46 |
+
) -> dict[str, Any]:
|
| 47 |
+
checkpoint = checkpoint.expanduser().resolve()
|
| 48 |
+
required = ("config.json", "model.safetensors")
|
| 49 |
+
missing = [name for name in required if not (checkpoint / name).is_file()]
|
| 50 |
+
if missing:
|
| 51 |
+
raise FileNotFoundError(
|
| 52 |
+
f"SmolVLA checkpoint is incomplete at {checkpoint}: missing {', '.join(missing)}"
|
| 53 |
+
)
|
| 54 |
+
|
| 55 |
+
result: dict[str, Any] = {
|
| 56 |
+
"schema_version": "smolvla-checkpoint-smoke/v0",
|
| 57 |
+
"checkpoint": str(checkpoint),
|
| 58 |
+
"metadata_only": metadata_only,
|
| 59 |
+
"required_files": list(required),
|
| 60 |
+
"package_versions": {
|
| 61 |
+
name: _package_version(name)
|
| 62 |
+
for name in ("lerobot", "torch", "transformers", "huggingface-hub")
|
| 63 |
+
},
|
| 64 |
+
}
|
| 65 |
+
if metadata_only:
|
| 66 |
+
result["ready"] = result["package_versions"]["lerobot"] is not None
|
| 67 |
+
return result
|
| 68 |
+
|
| 69 |
+
if vlm_metadata is None:
|
| 70 |
+
raise FileNotFoundError(
|
| 71 |
+
"Offline SmolVLA loading requires --vlm-metadata with local SmolVLM "
|
| 72 |
+
"config/tokenizer files"
|
| 73 |
+
)
|
| 74 |
+
vlm_metadata = vlm_metadata.expanduser().resolve()
|
| 75 |
+
vlm_required = ("config.json", "preprocessor_config.json", "tokenizer.json")
|
| 76 |
+
missing_vlm = [name for name in vlm_required if not (vlm_metadata / name).is_file()]
|
| 77 |
+
if missing_vlm:
|
| 78 |
+
raise FileNotFoundError(
|
| 79 |
+
f"SmolVLM metadata is incomplete at {vlm_metadata}: missing {', '.join(missing_vlm)}"
|
| 80 |
+
)
|
| 81 |
+
|
| 82 |
+
try:
|
| 83 |
+
import torch
|
| 84 |
+
|
| 85 |
+
from dovla_cil.eval.smolvla_runtime import (
|
| 86 |
+
import_smolvla_classes,
|
| 87 |
+
load_smolvla_config,
|
| 88 |
+
)
|
| 89 |
+
|
| 90 |
+
SmolVLAPolicy, _ = import_smolvla_classes()
|
| 91 |
+
except ImportError as exc:
|
| 92 |
+
raise ImportError(
|
| 93 |
+
'SmolVLA runtime is unavailable. Install isolated dependencies with '
|
| 94 |
+
'`pip install "lerobot[smolvla]==0.4.3"`. '
|
| 95 |
+
f"Original import error: {type(exc).__name__}: {exc}"
|
| 96 |
+
) from exc
|
| 97 |
+
|
| 98 |
+
resolved_device = "cuda" if device == "auto" and torch.cuda.is_available() else device
|
| 99 |
+
if resolved_device == "auto":
|
| 100 |
+
resolved_device = "cpu"
|
| 101 |
+
if resolved_device == "cuda" and not torch.cuda.is_available():
|
| 102 |
+
raise RuntimeError("CUDA was requested, but torch.cuda.is_available() is false")
|
| 103 |
+
|
| 104 |
+
print(json.dumps({"phase": "config_loading"}), flush=True)
|
| 105 |
+
config = load_smolvla_config(checkpoint, local_files_only=True)
|
| 106 |
+
config.device = resolved_device
|
| 107 |
+
config.vlm_model_name = str(vlm_metadata)
|
| 108 |
+
config.load_vlm_weights = False
|
| 109 |
+
|
| 110 |
+
print(json.dumps({"phase": "model_loading", "device": resolved_device}), flush=True)
|
| 111 |
+
started = time.perf_counter()
|
| 112 |
+
policy = SmolVLAPolicy.from_pretrained(
|
| 113 |
+
str(checkpoint),
|
| 114 |
+
config=config,
|
| 115 |
+
local_files_only=True,
|
| 116 |
+
)
|
| 117 |
+
policy = policy.to(torch.device(resolved_device)).eval()
|
| 118 |
+
load_seconds = time.perf_counter() - started
|
| 119 |
+
parameters = sum(parameter.numel() for parameter in policy.parameters())
|
| 120 |
+
trainable_parameters = sum(
|
| 121 |
+
parameter.numel() for parameter in policy.parameters() if parameter.requires_grad
|
| 122 |
+
)
|
| 123 |
+
print(json.dumps({"phase": "model_ready", "device": resolved_device}), flush=True)
|
| 124 |
+
result.update(
|
| 125 |
+
{
|
| 126 |
+
"ready": True,
|
| 127 |
+
"device": resolved_device,
|
| 128 |
+
"vlm_metadata": str(vlm_metadata),
|
| 129 |
+
"vlm_load_mode": "local_config_then_policy_safetensors",
|
| 130 |
+
"cuda_device": (
|
| 131 |
+
torch.cuda.get_device_name(0) if resolved_device == "cuda" else None
|
| 132 |
+
),
|
| 133 |
+
"load_seconds": load_seconds,
|
| 134 |
+
"parameter_count": parameters,
|
| 135 |
+
"trainable_parameter_count": trainable_parameters,
|
| 136 |
+
"policy_class": f"{type(policy).__module__}.{type(policy).__name__}",
|
| 137 |
+
}
|
| 138 |
+
)
|
| 139 |
+
return result
|
| 140 |
+
|
| 141 |
+
|
| 142 |
+
def _package_version(name: str) -> str | None:
|
| 143 |
+
try:
|
| 144 |
+
return importlib.metadata.version(name)
|
| 145 |
+
except importlib.metadata.PackageNotFoundError:
|
| 146 |
+
return None
|
| 147 |
+
|
| 148 |
+
|
| 149 |
+
def main() -> int:
|
| 150 |
+
args = build_parser().parse_args()
|
| 151 |
+
try:
|
| 152 |
+
result = smoke_checkpoint(
|
| 153 |
+
args.checkpoint,
|
| 154 |
+
device=args.device,
|
| 155 |
+
metadata_only=args.metadata_only,
|
| 156 |
+
vlm_metadata=args.vlm_metadata,
|
| 157 |
+
)
|
| 158 |
+
except (FileNotFoundError, ImportError, RuntimeError) as exc:
|
| 159 |
+
print(json.dumps({"ready": False, "error": str(exc)}, indent=2))
|
| 160 |
+
return 2
|
| 161 |
+
|
| 162 |
+
args.out.parent.mkdir(parents=True, exist_ok=True)
|
| 163 |
+
args.out.write_text(json.dumps(result, indent=2, sort_keys=True), encoding="utf-8")
|
| 164 |
+
print(json.dumps(result, indent=2, sort_keys=True))
|
| 165 |
+
return 0
|
| 166 |
+
|
| 167 |
+
|
| 168 |
+
if __name__ == "__main__":
|
| 169 |
+
raise SystemExit(main())
|
scripts/smoke_test.sh
ADDED
|
@@ -0,0 +1,27 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env bash
|
| 2 |
+
set -euo pipefail
|
| 3 |
+
|
| 4 |
+
ROOT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)"
|
| 5 |
+
cd "$ROOT_DIR"
|
| 6 |
+
|
| 7 |
+
export OPENCLAUDE_MOCK="${OPENCLAUDE_MOCK:-1}"
|
| 8 |
+
|
| 9 |
+
OUT_ROOT="${DOVLA_SMOKE_OUT:-outputs/phase5_smoke}"
|
| 10 |
+
TASKS_PATH="$OUT_ROOT/tasks.jsonl"
|
| 11 |
+
DATASET_DIR="$OUT_ROOT/cil"
|
| 12 |
+
|
| 13 |
+
mkdir -p "$OUT_ROOT"
|
| 14 |
+
|
| 15 |
+
python scripts/generate_tasks.py --mock --num-tasks 3 --out "$TASKS_PATH" --seed 0
|
| 16 |
+
python scripts/generate_cil.py \
|
| 17 |
+
--backend toy \
|
| 18 |
+
--tasks "$TASKS_PATH" \
|
| 19 |
+
--out "$DATASET_DIR" \
|
| 20 |
+
--num-states-per-task 2 \
|
| 21 |
+
--k 4 \
|
| 22 |
+
--seed 0 \
|
| 23 |
+
--shard-size 8 \
|
| 24 |
+
--inline-observations
|
| 25 |
+
python scripts/inspect_shard.py "$DATASET_DIR"
|
| 26 |
+
|
| 27 |
+
echo "smoke dataset: $DATASET_DIR"
|
scripts/train_dovla.py
ADDED
|
@@ -0,0 +1,153 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
from __future__ import annotations
|
| 3 |
+
|
| 4 |
+
import argparse
|
| 5 |
+
import sys
|
| 6 |
+
from dataclasses import fields
|
| 7 |
+
from pathlib import Path
|
| 8 |
+
|
| 9 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 10 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 11 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 12 |
+
|
| 13 |
+
from dovla_cil.training.losses import InterventionalLossWeights # noqa: E402
|
| 14 |
+
from dovla_cil.training.trainer import DoVLATrainer, TrainerConfig # noqa: E402
|
| 15 |
+
|
| 16 |
+
|
| 17 |
+
def main(argv: list[str] | None = None) -> int:
|
| 18 |
+
parser = argparse.ArgumentParser(description="Train the lightweight DoVLA model on CIL data.")
|
| 19 |
+
parser.add_argument("--dataset", type=Path, required=True)
|
| 20 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 21 |
+
parser.add_argument("--epochs", type=int, default=5)
|
| 22 |
+
parser.add_argument("--batch-groups", type=int, default=8)
|
| 23 |
+
parser.add_argument("--records-per-group", type=int, default=8)
|
| 24 |
+
parser.add_argument("--pair-count-per-group", type=int, default=8)
|
| 25 |
+
parser.add_argument("--hidden-dim", type=int, default=256)
|
| 26 |
+
parser.add_argument("--obs-dim", type=int, default=32)
|
| 27 |
+
parser.add_argument(
|
| 28 |
+
"--observation-mode",
|
| 29 |
+
choices=("state", "rgb"),
|
| 30 |
+
default="state",
|
| 31 |
+
help="Use inline state features or JPEG/HDF5 RGB observation references.",
|
| 32 |
+
)
|
| 33 |
+
parser.add_argument("--lang-dim", type=int, default=64)
|
| 34 |
+
parser.add_argument("--action-dim", type=int, default=8)
|
| 35 |
+
parser.add_argument("--action-horizon", type=int, default=4)
|
| 36 |
+
parser.add_argument("--effect-dim", type=int, default=32)
|
| 37 |
+
parser.add_argument(
|
| 38 |
+
"--backbone",
|
| 39 |
+
choices=("native", "clip"),
|
| 40 |
+
default="native",
|
| 41 |
+
help="Observation-language backbone. CLIP remains optional and locally loaded.",
|
| 42 |
+
)
|
| 43 |
+
parser.add_argument(
|
| 44 |
+
"--backbone-model",
|
| 45 |
+
help="Pinned local Hugging Face model directory for the optional CLIP backbone.",
|
| 46 |
+
)
|
| 47 |
+
parser.add_argument(
|
| 48 |
+
"--finetune-backbone",
|
| 49 |
+
action="store_true",
|
| 50 |
+
help="Fine-tune pretrained CLIP instead of the default frozen-feature regime.",
|
| 51 |
+
)
|
| 52 |
+
parser.add_argument(
|
| 53 |
+
"--backbone-feature-cache",
|
| 54 |
+
type=Path,
|
| 55 |
+
help="Reusable frozen CLIP feature cache shared across seeds.",
|
| 56 |
+
)
|
| 57 |
+
parser.add_argument("--backbone-feature-batch-size", type=int, default=64)
|
| 58 |
+
parser.add_argument("--lr", type=float, default=1e-3)
|
| 59 |
+
parser.add_argument("--weight-decay", type=float, default=0.0)
|
| 60 |
+
parser.add_argument("--device", default="auto")
|
| 61 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 62 |
+
parser.add_argument("--val-fraction", type=float, default=0.2)
|
| 63 |
+
parser.add_argument("--wandb", action="store_true", help="Enable wandb if installed.")
|
| 64 |
+
parser.add_argument(
|
| 65 |
+
"--objective",
|
| 66 |
+
choices=("lattice_field", "legacy"),
|
| 67 |
+
default="lattice_field",
|
| 68 |
+
help="Use the proposed interventional field objective or the legacy multi-head ablation.",
|
| 69 |
+
)
|
| 70 |
+
parser.add_argument(
|
| 71 |
+
"--lattice-neighbors",
|
| 72 |
+
type=int,
|
| 73 |
+
default=32,
|
| 74 |
+
help="Nearest action neighbors per node; 32 gives a complete graph for K<=32.",
|
| 75 |
+
)
|
| 76 |
+
parser.add_argument(
|
| 77 |
+
"--pair-scope",
|
| 78 |
+
choices=("same_state", "cross_state"),
|
| 79 |
+
default="same_state",
|
| 80 |
+
help="Choose whether legacy ranking pairs share the exact simulator state.",
|
| 81 |
+
)
|
| 82 |
+
parser.add_argument(
|
| 83 |
+
"--loss-weight",
|
| 84 |
+
action="append",
|
| 85 |
+
default=[],
|
| 86 |
+
metavar="NAME=VALUE",
|
| 87 |
+
help="Override one loss weight; repeat for multiple weights.",
|
| 88 |
+
)
|
| 89 |
+
args = parser.parse_args(argv)
|
| 90 |
+
try:
|
| 91 |
+
loss_weights = _parse_loss_weights(args.loss_weight)
|
| 92 |
+
except ValueError as exc:
|
| 93 |
+
parser.error(str(exc))
|
| 94 |
+
|
| 95 |
+
config = TrainerConfig(
|
| 96 |
+
dataset_dir=args.dataset,
|
| 97 |
+
output_dir=args.out,
|
| 98 |
+
epochs=args.epochs,
|
| 99 |
+
batch_groups=args.batch_groups,
|
| 100 |
+
records_per_group=args.records_per_group,
|
| 101 |
+
pair_count_per_group=args.pair_count_per_group,
|
| 102 |
+
hidden_dim=args.hidden_dim,
|
| 103 |
+
obs_dim=args.obs_dim,
|
| 104 |
+
observation_mode=args.observation_mode,
|
| 105 |
+
lang_dim=args.lang_dim,
|
| 106 |
+
action_dim=args.action_dim,
|
| 107 |
+
action_horizon=args.action_horizon,
|
| 108 |
+
effect_dim=args.effect_dim,
|
| 109 |
+
backbone_type=args.backbone,
|
| 110 |
+
backbone_model=args.backbone_model,
|
| 111 |
+
backbone_freeze=not args.finetune_backbone,
|
| 112 |
+
backbone_feature_cache=args.backbone_feature_cache,
|
| 113 |
+
backbone_feature_batch_size=args.backbone_feature_batch_size,
|
| 114 |
+
learning_rate=args.lr,
|
| 115 |
+
weight_decay=args.weight_decay,
|
| 116 |
+
device=args.device,
|
| 117 |
+
seed=args.seed,
|
| 118 |
+
val_fraction=args.val_fraction,
|
| 119 |
+
wandb=args.wandb,
|
| 120 |
+
objective=args.objective,
|
| 121 |
+
lattice_neighbors=args.lattice_neighbors,
|
| 122 |
+
pair_scope=args.pair_scope,
|
| 123 |
+
losses=loss_weights,
|
| 124 |
+
)
|
| 125 |
+
result = DoVLATrainer(config).train()
|
| 126 |
+
best = result.get("best", {})
|
| 127 |
+
print(f"wrote checkpoints to {args.out}")
|
| 128 |
+
print(f"best val rank_acc={best.get('rank_acc', 0.0):.4f}")
|
| 129 |
+
return 0
|
| 130 |
+
|
| 131 |
+
|
| 132 |
+
def _parse_loss_weights(items: list[str]) -> InterventionalLossWeights:
|
| 133 |
+
allowed = {field.name for field in fields(InterventionalLossWeights)}
|
| 134 |
+
values: dict[str, float] = {}
|
| 135 |
+
for item in items:
|
| 136 |
+
if "=" not in item:
|
| 137 |
+
raise ValueError(f"loss weight must use NAME=VALUE syntax: {item!r}")
|
| 138 |
+
name, raw_value = item.split("=", 1)
|
| 139 |
+
if name not in allowed:
|
| 140 |
+
choices = ", ".join(sorted(allowed))
|
| 141 |
+
raise ValueError(f"unknown loss weight {name!r}; choose one of: {choices}")
|
| 142 |
+
try:
|
| 143 |
+
value = float(raw_value)
|
| 144 |
+
except ValueError as exc:
|
| 145 |
+
raise ValueError(f"loss weight {name!r} must be numeric") from exc
|
| 146 |
+
if value < 0:
|
| 147 |
+
raise ValueError(f"loss weight {name!r} must be non-negative")
|
| 148 |
+
values[name] = value
|
| 149 |
+
return InterventionalLossWeights(**values)
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
if __name__ == "__main__":
|
| 153 |
+
raise SystemExit(main())
|
scripts/train_dovla_attention.py
ADDED
|
@@ -0,0 +1,324 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""
|
| 3 |
+
Standalone trainer for DoVLA-Attention (CVPR submission)
|
| 4 |
+
|
| 5 |
+
Single architectural contribution: Transformer attention for action comparison
|
| 6 |
+
- Cross-attention: observation conditions actions
|
| 7 |
+
- Self-attention: models action relationships
|
| 8 |
+
- Pairwise comparison: structured features
|
| 9 |
+
|
| 10 |
+
Expected: 42-44% success (vs 38.43% MLP baseline)
|
| 11 |
+
"""
|
| 12 |
+
from __future__ import annotations
|
| 13 |
+
|
| 14 |
+
import argparse
|
| 15 |
+
import json
|
| 16 |
+
import random
|
| 17 |
+
import sys
|
| 18 |
+
from pathlib import Path
|
| 19 |
+
from typing import Optional
|
| 20 |
+
|
| 21 |
+
import numpy as np
|
| 22 |
+
import torch
|
| 23 |
+
import torch.nn as nn
|
| 24 |
+
import torch.optim as optim
|
| 25 |
+
from torch.utils.data import DataLoader, Dataset
|
| 26 |
+
|
| 27 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 28 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 29 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 30 |
+
|
| 31 |
+
from dovla_cil.models.dovla_attention import DoVLAAttention
|
| 32 |
+
from dovla_cil.data.cil_collection import CILCollection
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
class AttentionTrainingDataset(Dataset):
|
| 36 |
+
"""Dataset for training DoVLA-Attention with pairwise ranking."""
|
| 37 |
+
|
| 38 |
+
def __init__(self, collection: CILCollection, group_ids: list[str],
|
| 39 |
+
records_per_group: int = 8, pairs_per_group: int = 8):
|
| 40 |
+
self.collection = collection
|
| 41 |
+
self.group_ids = group_ids
|
| 42 |
+
self.records_per_group = records_per_group
|
| 43 |
+
self.pairs_per_group = pairs_per_group
|
| 44 |
+
|
| 45 |
+
def __len__(self):
|
| 46 |
+
return len(self.group_ids)
|
| 47 |
+
|
| 48 |
+
def __getitem__(self, idx):
|
| 49 |
+
group_id = self.group_ids[idx]
|
| 50 |
+
records = self.collection.get_group_records(group_id)
|
| 51 |
+
|
| 52 |
+
# Sample records
|
| 53 |
+
if len(records) > self.records_per_group:
|
| 54 |
+
records = random.sample(records, self.records_per_group)
|
| 55 |
+
|
| 56 |
+
# Extract data
|
| 57 |
+
obs = records[0]["observation"] # Same state for all
|
| 58 |
+
actions = [r["action"] for r in records]
|
| 59 |
+
rewards = [r.get("reward", 0.0) for r in records]
|
| 60 |
+
|
| 61 |
+
# Sample pairs for ranking loss
|
| 62 |
+
pairs = []
|
| 63 |
+
for _ in range(self.pairs_per_group):
|
| 64 |
+
i, j = random.sample(range(len(records)), 2)
|
| 65 |
+
if rewards[i] != rewards[j]:
|
| 66 |
+
pairs.append((i, j, 1.0 if rewards[i] > rewards[j] else 0.0))
|
| 67 |
+
|
| 68 |
+
return {
|
| 69 |
+
"observation": torch.FloatTensor(obs),
|
| 70 |
+
"actions": torch.FloatTensor(actions),
|
| 71 |
+
"rewards": torch.FloatTensor(rewards),
|
| 72 |
+
"pairs": pairs
|
| 73 |
+
}
|
| 74 |
+
|
| 75 |
+
|
| 76 |
+
def set_seed(seed: int):
|
| 77 |
+
"""Set all random seeds for reproducibility."""
|
| 78 |
+
random.seed(seed)
|
| 79 |
+
np.random.seed(seed)
|
| 80 |
+
torch.manual_seed(seed)
|
| 81 |
+
if torch.cuda.is_available():
|
| 82 |
+
torch.cuda.manual_seed_all(seed)
|
| 83 |
+
|
| 84 |
+
|
| 85 |
+
def train_epoch(model: nn.Module, dataloader: DataLoader,
|
| 86 |
+
optimizer: optim.Optimizer, device: torch.device) -> float:
|
| 87 |
+
"""Train for one epoch."""
|
| 88 |
+
model.train()
|
| 89 |
+
total_loss = 0.0
|
| 90 |
+
num_batches = 0
|
| 91 |
+
|
| 92 |
+
for batch in dataloader:
|
| 93 |
+
obs = batch["observation"].to(device)
|
| 94 |
+
actions = batch["actions"].to(device)
|
| 95 |
+
rewards = batch["rewards"].to(device)
|
| 96 |
+
|
| 97 |
+
optimizer.zero_grad()
|
| 98 |
+
|
| 99 |
+
# Forward: get pairwise scores
|
| 100 |
+
scores = model(obs, actions) # (batch, K, K)
|
| 101 |
+
|
| 102 |
+
# Ranking loss: prefer higher reward actions
|
| 103 |
+
batch_size, K = rewards.shape
|
| 104 |
+
loss = 0.0
|
| 105 |
+
count = 0
|
| 106 |
+
|
| 107 |
+
for b in range(batch_size):
|
| 108 |
+
for i in range(K):
|
| 109 |
+
for j in range(K):
|
| 110 |
+
if i != j and rewards[b, i] != rewards[b, j]:
|
| 111 |
+
target = 1.0 if rewards[b, i] > rewards[b, j] else 0.0
|
| 112 |
+
pred = torch.sigmoid(scores[b, i, j])
|
| 113 |
+
loss += nn.functional.binary_cross_entropy(
|
| 114 |
+
pred.unsqueeze(0),
|
| 115 |
+
torch.tensor([target], device=device)
|
| 116 |
+
)
|
| 117 |
+
count += 1
|
| 118 |
+
|
| 119 |
+
if count > 0:
|
| 120 |
+
loss = loss / count
|
| 121 |
+
loss.backward()
|
| 122 |
+
optimizer.step()
|
| 123 |
+
|
| 124 |
+
total_loss += loss.item()
|
| 125 |
+
num_batches += 1
|
| 126 |
+
|
| 127 |
+
return total_loss / max(num_batches, 1)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def evaluate(model: nn.Module, dataloader: DataLoader, device: torch.device) -> dict:
|
| 131 |
+
"""Evaluate model on validation set."""
|
| 132 |
+
model.eval()
|
| 133 |
+
correct = 0
|
| 134 |
+
total = 0
|
| 135 |
+
|
| 136 |
+
with torch.no_grad():
|
| 137 |
+
for batch in dataloader:
|
| 138 |
+
obs = batch["observation"].to(device)
|
| 139 |
+
actions = batch["actions"].to(device)
|
| 140 |
+
rewards = batch["rewards"].to(device)
|
| 141 |
+
|
| 142 |
+
scores = model(obs, actions)
|
| 143 |
+
|
| 144 |
+
batch_size, K = rewards.shape
|
| 145 |
+
for b in range(batch_size):
|
| 146 |
+
for i in range(K):
|
| 147 |
+
for j in range(K):
|
| 148 |
+
if i != j and rewards[b, i] != rewards[b, j]:
|
| 149 |
+
pred = scores[b, i, j] > 0
|
| 150 |
+
target = rewards[b, i] > rewards[b, j]
|
| 151 |
+
if pred == target:
|
| 152 |
+
correct += 1
|
| 153 |
+
total += 1
|
| 154 |
+
|
| 155 |
+
return {
|
| 156 |
+
"accuracy": correct / max(total, 1),
|
| 157 |
+
"correct": correct,
|
| 158 |
+
"total": total
|
| 159 |
+
}
|
| 160 |
+
|
| 161 |
+
|
| 162 |
+
def main(argv: list[str] | None = None) -> int:
|
| 163 |
+
parser = argparse.ArgumentParser(
|
| 164 |
+
description="Train DoVLA-Attention for CVPR paper"
|
| 165 |
+
)
|
| 166 |
+
|
| 167 |
+
# Data
|
| 168 |
+
parser.add_argument("--dataset", type=Path, required=True)
|
| 169 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 170 |
+
|
| 171 |
+
# Architecture
|
| 172 |
+
parser.add_argument("--hidden-dim", type=int, default=256)
|
| 173 |
+
parser.add_argument("--n-heads", type=int, default=4)
|
| 174 |
+
parser.add_argument("--n-layers", type=int, default=2)
|
| 175 |
+
|
| 176 |
+
# Training
|
| 177 |
+
parser.add_argument("--epochs", type=int, default=50)
|
| 178 |
+
parser.add_argument("--batch-size", type=int, default=16)
|
| 179 |
+
parser.add_argument("--lr", type=float, default=0.0003)
|
| 180 |
+
parser.add_argument("--weight-decay", type=float, default=0.01)
|
| 181 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 182 |
+
parser.add_argument("--val-fraction", type=float, default=0.2)
|
| 183 |
+
|
| 184 |
+
# System
|
| 185 |
+
parser.add_argument("--device", default="auto")
|
| 186 |
+
|
| 187 |
+
args = parser.parse_args(argv)
|
| 188 |
+
|
| 189 |
+
# Setup
|
| 190 |
+
set_seed(args.seed)
|
| 191 |
+
args.out.mkdir(parents=True, exist_ok=True)
|
| 192 |
+
|
| 193 |
+
if args.device == "auto":
|
| 194 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 195 |
+
else:
|
| 196 |
+
device = torch.device(args.device)
|
| 197 |
+
|
| 198 |
+
print("=" * 70)
|
| 199 |
+
print("DoVLA-Attention Training (CVPR)")
|
| 200 |
+
print("=" * 70)
|
| 201 |
+
print(f"Dataset: {args.dataset}")
|
| 202 |
+
print(f"Output: {args.out}")
|
| 203 |
+
print(f"Device: {device}")
|
| 204 |
+
print(f"Hidden dim: {args.hidden_dim}")
|
| 205 |
+
print(f"Heads: {args.n_heads}, Layers: {args.n_layers}")
|
| 206 |
+
print(f"Seed: {args.seed}")
|
| 207 |
+
print()
|
| 208 |
+
|
| 209 |
+
# Load data
|
| 210 |
+
print("Loading dataset...")
|
| 211 |
+
collection = CILCollection(args.dataset)
|
| 212 |
+
all_groups = list(collection.group_ids)
|
| 213 |
+
|
| 214 |
+
# Split train/val
|
| 215 |
+
random.shuffle(all_groups)
|
| 216 |
+
split_idx = int(len(all_groups) * (1 - args.val_fraction))
|
| 217 |
+
train_groups = all_groups[:split_idx]
|
| 218 |
+
val_groups = all_groups[split_idx:]
|
| 219 |
+
|
| 220 |
+
print(f"Total groups: {len(all_groups)}")
|
| 221 |
+
print(f"Train: {len(train_groups)}, Val: {len(val_groups)}")
|
| 222 |
+
print()
|
| 223 |
+
|
| 224 |
+
# Create datasets
|
| 225 |
+
train_dataset = AttentionTrainingDataset(collection, train_groups)
|
| 226 |
+
val_dataset = AttentionTrainingDataset(collection, val_groups)
|
| 227 |
+
|
| 228 |
+
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
|
| 229 |
+
shuffle=True, num_workers=0)
|
| 230 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
|
| 231 |
+
shuffle=False, num_workers=0)
|
| 232 |
+
|
| 233 |
+
# Get dims from first sample
|
| 234 |
+
sample = train_dataset[0]
|
| 235 |
+
obs_dim = sample["observation"].shape[0]
|
| 236 |
+
action_dim = sample["actions"].shape[1]
|
| 237 |
+
|
| 238 |
+
print(f"Observation dim: {obs_dim}")
|
| 239 |
+
print(f"Action dim: {action_dim}")
|
| 240 |
+
print()
|
| 241 |
+
|
| 242 |
+
# Create model
|
| 243 |
+
model = DoVLAAttention(
|
| 244 |
+
obs_dim=obs_dim,
|
| 245 |
+
action_dim=action_dim,
|
| 246 |
+
hidden_dim=args.hidden_dim,
|
| 247 |
+
n_heads=args.n_heads,
|
| 248 |
+
n_layers=args.n_layers
|
| 249 |
+
).to(device)
|
| 250 |
+
|
| 251 |
+
num_params = sum(p.numel() for p in model.parameters())
|
| 252 |
+
print(f"Model parameters: {num_params:,}")
|
| 253 |
+
print()
|
| 254 |
+
|
| 255 |
+
# Optimizer
|
| 256 |
+
optimizer = optim.AdamW(model.parameters(), lr=args.lr,
|
| 257 |
+
weight_decay=args.weight_decay)
|
| 258 |
+
|
| 259 |
+
# Training loop
|
| 260 |
+
best_acc = 0.0
|
| 261 |
+
history = []
|
| 262 |
+
|
| 263 |
+
print("Starting training...")
|
| 264 |
+
print()
|
| 265 |
+
|
| 266 |
+
for epoch in range(args.epochs):
|
| 267 |
+
train_loss = train_epoch(model, train_loader, optimizer, device)
|
| 268 |
+
val_metrics = evaluate(model, val_loader, device)
|
| 269 |
+
|
| 270 |
+
val_acc = val_metrics["accuracy"]
|
| 271 |
+
|
| 272 |
+
history.append({
|
| 273 |
+
"epoch": epoch + 1,
|
| 274 |
+
"train_loss": train_loss,
|
| 275 |
+
"val_accuracy": val_acc
|
| 276 |
+
})
|
| 277 |
+
|
| 278 |
+
print(f"Epoch {epoch+1:3d}/{args.epochs}: "
|
| 279 |
+
f"loss={train_loss:.4f}, val_acc={val_acc:.4f}")
|
| 280 |
+
|
| 281 |
+
# Save best model
|
| 282 |
+
if val_acc > best_acc:
|
| 283 |
+
best_acc = val_acc
|
| 284 |
+
torch.save({
|
| 285 |
+
"model_state_dict": model.state_dict(),
|
| 286 |
+
"epoch": epoch + 1,
|
| 287 |
+
"val_accuracy": val_acc,
|
| 288 |
+
"args": vars(args)
|
| 289 |
+
}, args.out / "best.pt")
|
| 290 |
+
|
| 291 |
+
print()
|
| 292 |
+
print(f"✅ Training complete! Best val accuracy: {best_acc:.4f}")
|
| 293 |
+
|
| 294 |
+
# Save training history
|
| 295 |
+
with open(args.out / "history.json", "w") as f:
|
| 296 |
+
json.dump(history, f, indent=2)
|
| 297 |
+
|
| 298 |
+
# Save final config
|
| 299 |
+
with open(args.out / "config.json", "w") as f:
|
| 300 |
+
json.dump({
|
| 301 |
+
"model": "DoVLA-Attention",
|
| 302 |
+
"architecture": {
|
| 303 |
+
"hidden_dim": args.hidden_dim,
|
| 304 |
+
"n_heads": args.n_heads,
|
| 305 |
+
"n_layers": args.n_layers
|
| 306 |
+
},
|
| 307 |
+
"training": {
|
| 308 |
+
"epochs": args.epochs,
|
| 309 |
+
"lr": args.lr,
|
| 310 |
+
"weight_decay": args.weight_decay,
|
| 311 |
+
"seed": args.seed
|
| 312 |
+
},
|
| 313 |
+
"results": {
|
| 314 |
+
"best_val_accuracy": best_acc,
|
| 315 |
+
"num_parameters": num_params
|
| 316 |
+
}
|
| 317 |
+
}, f, indent=2)
|
| 318 |
+
|
| 319 |
+
print(f"Saved to: {args.out}")
|
| 320 |
+
return 0
|
| 321 |
+
|
| 322 |
+
|
| 323 |
+
if __name__ == "__main__":
|
| 324 |
+
sys.exit(main())
|
scripts/train_dovla_enhanced.py
ADDED
|
@@ -0,0 +1,407 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""
|
| 3 |
+
Enhanced DoVLA-Attention Trainer with SOTA Components
|
| 4 |
+
|
| 5 |
+
Architecture improvements:
|
| 6 |
+
1. Hierarchical attention (local + global)
|
| 7 |
+
2. Graph neural networks (explicit structure)
|
| 8 |
+
3. Contrastive learning (better embeddings)
|
| 9 |
+
4. Task-adaptive layers (multi-task)
|
| 10 |
+
|
| 11 |
+
Expected: 44-47% success (vs 38.43% baseline, +5.5-8.5%)
|
| 12 |
+
"""
|
| 13 |
+
from __future__ import annotations
|
| 14 |
+
|
| 15 |
+
import argparse
|
| 16 |
+
import json
|
| 17 |
+
import random
|
| 18 |
+
import sys
|
| 19 |
+
from pathlib import Path
|
| 20 |
+
from typing import Optional
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
import torch
|
| 24 |
+
import torch.nn as nn
|
| 25 |
+
import torch.optim as optim
|
| 26 |
+
from torch.utils.data import DataLoader, Dataset
|
| 27 |
+
|
| 28 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 29 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 30 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 31 |
+
|
| 32 |
+
from dovla_cil.models.dovla_attention_enhanced import DoVLAAttentionEnhanced
|
| 33 |
+
from dovla_cil.data.datasets import CILDataset
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class EnhancedTrainingDataset(Dataset):
|
| 37 |
+
"""Dataset for enhanced architecture with task IDs and rewards."""
|
| 38 |
+
|
| 39 |
+
def __init__(self, dataset: CILDataset, group_ids: list[str],
|
| 40 |
+
records_per_group: int = 16,
|
| 41 |
+
max_obs_dim: int = 70, max_act_dim: int = 32):
|
| 42 |
+
self.dataset = dataset
|
| 43 |
+
self.group_ids = group_ids
|
| 44 |
+
self.records_per_group = records_per_group
|
| 45 |
+
# Pad all observations/actions to fixed max dims for multi-task batching.
|
| 46 |
+
# This is a standard, fair approach: every method sees the same padded space.
|
| 47 |
+
self.max_obs_dim = max_obs_dim
|
| 48 |
+
self.max_act_dim = max_act_dim
|
| 49 |
+
|
| 50 |
+
# Task name to ID mapping
|
| 51 |
+
self.task_map = {
|
| 52 |
+
"PickCube-v1": 0,
|
| 53 |
+
"PushCube-v1": 1,
|
| 54 |
+
"PullCube-v1": 2,
|
| 55 |
+
"StackCube-v1": 3,
|
| 56 |
+
"LiftPegUpright-v1": 4,
|
| 57 |
+
"PegInsertionSide-v1": 5
|
| 58 |
+
}
|
| 59 |
+
|
| 60 |
+
def _pad(self, vec: list[float], target: int) -> list[float]:
|
| 61 |
+
if len(vec) >= target:
|
| 62 |
+
return vec[:target]
|
| 63 |
+
return vec + [0.0] * (target - len(vec))
|
| 64 |
+
|
| 65 |
+
def __len__(self):
|
| 66 |
+
return len(self.group_ids)
|
| 67 |
+
|
| 68 |
+
def __getitem__(self, idx):
|
| 69 |
+
group_id = self.group_ids[idx]
|
| 70 |
+
records = self.dataset.get_group(group_id)
|
| 71 |
+
|
| 72 |
+
# Sample more records for better training
|
| 73 |
+
if len(records) > self.records_per_group:
|
| 74 |
+
records = random.sample(records, self.records_per_group)
|
| 75 |
+
|
| 76 |
+
# Extract data from CILRecord objects
|
| 77 |
+
# Observation can be inline dict or reference (use inline if available)
|
| 78 |
+
obs_data = records[0].observation_inline
|
| 79 |
+
if obs_data is None or not isinstance(obs_data, dict):
|
| 80 |
+
raise ValueError(f"No inline observation for group {group_id}")
|
| 81 |
+
|
| 82 |
+
# Convert observation dict to flat array
|
| 83 |
+
if "state" in obs_data:
|
| 84 |
+
obs = list(obs_data["state"])
|
| 85 |
+
else:
|
| 86 |
+
# Flatten all numeric values
|
| 87 |
+
obs = []
|
| 88 |
+
for v in obs_data.values():
|
| 89 |
+
if isinstance(v, list):
|
| 90 |
+
obs.extend(v)
|
| 91 |
+
elif isinstance(v, (int, float)):
|
| 92 |
+
obs.append(v)
|
| 93 |
+
|
| 94 |
+
# Pad observation to fixed dim for multi-task batching
|
| 95 |
+
obs = self._pad([float(x) for x in obs], self.max_obs_dim)
|
| 96 |
+
|
| 97 |
+
# Extract actions from action_chunk, pad to fixed dim
|
| 98 |
+
actions = [self._pad(r.action_chunk.flat_values, self.max_act_dim) for r in records]
|
| 99 |
+
|
| 100 |
+
# Extract rewards
|
| 101 |
+
rewards = [r.reward.score for r in records]
|
| 102 |
+
|
| 103 |
+
# Get task ID
|
| 104 |
+
task_name = records[0].task_id
|
| 105 |
+
task_id = self.task_map.get(task_name, 0)
|
| 106 |
+
|
| 107 |
+
return {
|
| 108 |
+
"observation": torch.FloatTensor(obs),
|
| 109 |
+
"actions": torch.FloatTensor(actions),
|
| 110 |
+
"rewards": torch.FloatTensor(rewards),
|
| 111 |
+
"task_id": torch.LongTensor([task_id])
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def collate_fn(batch):
|
| 116 |
+
"""Custom collate to handle variable K and ensure fixed dims."""
|
| 117 |
+
max_k = max(b["actions"].shape[0] for b in batch)
|
| 118 |
+
|
| 119 |
+
batch_size = len(batch)
|
| 120 |
+
obs_dim = batch[0]["observation"].shape[0] # Already padded to fixed dim
|
| 121 |
+
action_dim = batch[0]["actions"].shape[1] # Already padded to fixed dim
|
| 122 |
+
|
| 123 |
+
# Stack observations (all same size after padding)
|
| 124 |
+
obs_batch = torch.stack([b["observation"] for b in batch])
|
| 125 |
+
|
| 126 |
+
# Pad actions to max K in batch
|
| 127 |
+
actions_batch = torch.zeros(batch_size, max_k, action_dim)
|
| 128 |
+
rewards_batch = torch.zeros(batch_size, max_k)
|
| 129 |
+
task_ids = torch.cat([b["task_id"] for b in batch])
|
| 130 |
+
|
| 131 |
+
for i, b in enumerate(batch):
|
| 132 |
+
k = b["actions"].shape[0]
|
| 133 |
+
actions_batch[i, :k] = b["actions"]
|
| 134 |
+
rewards_batch[i, :k] = b["rewards"]
|
| 135 |
+
|
| 136 |
+
return {
|
| 137 |
+
"observation": obs_batch,
|
| 138 |
+
"actions": actions_batch,
|
| 139 |
+
"rewards": rewards_batch,
|
| 140 |
+
"task_id": task_ids
|
| 141 |
+
}
|
| 142 |
+
|
| 143 |
+
|
| 144 |
+
def set_seed(seed: int):
|
| 145 |
+
random.seed(seed)
|
| 146 |
+
np.random.seed(seed)
|
| 147 |
+
torch.manual_seed(seed)
|
| 148 |
+
if torch.cuda.is_available():
|
| 149 |
+
torch.cuda.manual_seed_all(seed)
|
| 150 |
+
|
| 151 |
+
|
| 152 |
+
def train_epoch(model: nn.Module, dataloader: DataLoader,
|
| 153 |
+
optimizer: optim.Optimizer, device: torch.device,
|
| 154 |
+
contrastive_weight: float = 0.1) -> dict:
|
| 155 |
+
"""Train for one epoch with ranking + contrastive losses."""
|
| 156 |
+
model.train()
|
| 157 |
+
total_ranking_loss = 0.0
|
| 158 |
+
total_contrastive_loss = 0.0
|
| 159 |
+
num_batches = 0
|
| 160 |
+
|
| 161 |
+
for batch in dataloader:
|
| 162 |
+
obs = batch["observation"].to(device)
|
| 163 |
+
actions = batch["actions"].to(device)
|
| 164 |
+
rewards = batch["rewards"].to(device)
|
| 165 |
+
task_ids = batch["task_id"].to(device).squeeze()
|
| 166 |
+
|
| 167 |
+
optimizer.zero_grad()
|
| 168 |
+
|
| 169 |
+
# Forward pass
|
| 170 |
+
scores, contrastive_loss = model(obs, actions, task_ids, rewards)
|
| 171 |
+
|
| 172 |
+
# Ranking loss
|
| 173 |
+
batch_size, K = rewards.shape
|
| 174 |
+
ranking_loss = 0.0
|
| 175 |
+
count = 0
|
| 176 |
+
|
| 177 |
+
for b in range(batch_size):
|
| 178 |
+
for i in range(K):
|
| 179 |
+
for j in range(K):
|
| 180 |
+
if i != j and rewards[b, i] != rewards[b, j]:
|
| 181 |
+
target = 1.0 if rewards[b, i] > rewards[b, j] else 0.0
|
| 182 |
+
pred = torch.sigmoid(scores[b, i, j])
|
| 183 |
+
ranking_loss += nn.functional.binary_cross_entropy(
|
| 184 |
+
pred.unsqueeze(0),
|
| 185 |
+
torch.tensor([target], device=device)
|
| 186 |
+
)
|
| 187 |
+
count += 1
|
| 188 |
+
|
| 189 |
+
if count > 0:
|
| 190 |
+
ranking_loss = ranking_loss / count
|
| 191 |
+
|
| 192 |
+
# Total loss
|
| 193 |
+
total_loss = ranking_loss
|
| 194 |
+
if contrastive_loss is not None:
|
| 195 |
+
total_loss = total_loss + contrastive_weight * contrastive_loss
|
| 196 |
+
|
| 197 |
+
total_loss.backward()
|
| 198 |
+
|
| 199 |
+
# Gradient clipping for stability
|
| 200 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 201 |
+
|
| 202 |
+
optimizer.step()
|
| 203 |
+
|
| 204 |
+
total_ranking_loss += ranking_loss.item()
|
| 205 |
+
if contrastive_loss is not None:
|
| 206 |
+
total_contrastive_loss += contrastive_loss.item()
|
| 207 |
+
num_batches += 1
|
| 208 |
+
|
| 209 |
+
return {
|
| 210 |
+
"ranking_loss": total_ranking_loss / max(num_batches, 1),
|
| 211 |
+
"contrastive_loss": total_contrastive_loss / max(num_batches, 1)
|
| 212 |
+
}
|
| 213 |
+
|
| 214 |
+
|
| 215 |
+
def evaluate(model: nn.Module, dataloader: DataLoader, device: torch.device) -> dict:
|
| 216 |
+
"""Evaluate model."""
|
| 217 |
+
model.eval()
|
| 218 |
+
correct = 0
|
| 219 |
+
total = 0
|
| 220 |
+
|
| 221 |
+
with torch.no_grad():
|
| 222 |
+
for batch in dataloader:
|
| 223 |
+
obs = batch["observation"].to(device)
|
| 224 |
+
actions = batch["actions"].to(device)
|
| 225 |
+
rewards = batch["rewards"].to(device)
|
| 226 |
+
task_ids = batch["task_id"].to(device).squeeze()
|
| 227 |
+
|
| 228 |
+
scores, _ = model(obs, actions, task_ids, None)
|
| 229 |
+
|
| 230 |
+
batch_size, K = rewards.shape
|
| 231 |
+
for b in range(batch_size):
|
| 232 |
+
for i in range(K):
|
| 233 |
+
for j in range(K):
|
| 234 |
+
if i != j and rewards[b, i] != rewards[b, j]:
|
| 235 |
+
pred = scores[b, i, j] > 0
|
| 236 |
+
target = rewards[b, i] > rewards[b, j]
|
| 237 |
+
if pred == target:
|
| 238 |
+
correct += 1
|
| 239 |
+
total += 1
|
| 240 |
+
|
| 241 |
+
return {
|
| 242 |
+
"accuracy": correct / max(total, 1)
|
| 243 |
+
}
|
| 244 |
+
|
| 245 |
+
|
| 246 |
+
def main(argv: list[str] | None = None) -> int:
|
| 247 |
+
parser = argparse.ArgumentParser(
|
| 248 |
+
description="Train Enhanced DoVLA-Attention for CVPR"
|
| 249 |
+
)
|
| 250 |
+
|
| 251 |
+
# Data
|
| 252 |
+
parser.add_argument("--dataset", type=Path, required=True)
|
| 253 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 254 |
+
|
| 255 |
+
# Architecture
|
| 256 |
+
parser.add_argument("--hidden-dim", type=int, default=256)
|
| 257 |
+
parser.add_argument("--n-heads", type=int, default=4)
|
| 258 |
+
parser.add_argument("--n-layers", type=int, default=3)
|
| 259 |
+
|
| 260 |
+
# Training
|
| 261 |
+
parser.add_argument("--epochs", type=int, default=50)
|
| 262 |
+
parser.add_argument("--batch-size", type=int, default=16)
|
| 263 |
+
parser.add_argument("--lr", type=float, default=0.0003)
|
| 264 |
+
parser.add_argument("--weight-decay", type=float, default=0.01)
|
| 265 |
+
parser.add_argument("--contrastive-weight", type=float, default=0.1)
|
| 266 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 267 |
+
parser.add_argument("--val-fraction", type=float, default=0.2)
|
| 268 |
+
|
| 269 |
+
# System
|
| 270 |
+
parser.add_argument("--device", default="auto")
|
| 271 |
+
|
| 272 |
+
args = parser.parse_args(argv)
|
| 273 |
+
|
| 274 |
+
set_seed(args.seed)
|
| 275 |
+
args.out.mkdir(parents=True, exist_ok=True)
|
| 276 |
+
|
| 277 |
+
if args.device == "auto":
|
| 278 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 279 |
+
else:
|
| 280 |
+
device = torch.device(args.device)
|
| 281 |
+
|
| 282 |
+
print("=" * 70)
|
| 283 |
+
print("Enhanced DoVLA-Attention Training (CVPR)")
|
| 284 |
+
print("=" * 70)
|
| 285 |
+
print(f"Dataset: {args.dataset}")
|
| 286 |
+
print(f"Device: {device}")
|
| 287 |
+
print(f"Architecture: Hierarchical + Graph + Contrastive + Task-Adaptive")
|
| 288 |
+
print(f"Hidden: {args.hidden_dim}, Heads: {args.n_heads}, Layers: {args.n_layers}")
|
| 289 |
+
print(f"Seed: {args.seed}")
|
| 290 |
+
print()
|
| 291 |
+
|
| 292 |
+
# Load data
|
| 293 |
+
print("Loading dataset...")
|
| 294 |
+
dataset = CILDataset(args.dataset)
|
| 295 |
+
all_groups = list(dataset.group_ids)
|
| 296 |
+
|
| 297 |
+
random.shuffle(all_groups)
|
| 298 |
+
split_idx = int(len(all_groups) * (1 - args.val_fraction))
|
| 299 |
+
train_groups = all_groups[:split_idx]
|
| 300 |
+
val_groups = all_groups[split_idx:]
|
| 301 |
+
|
| 302 |
+
print(f"Total: {len(all_groups)}, Train: {len(train_groups)}, Val: {len(val_groups)}")
|
| 303 |
+
print()
|
| 304 |
+
|
| 305 |
+
# Datasets
|
| 306 |
+
train_dataset = EnhancedTrainingDataset(dataset, train_groups, records_per_group=16)
|
| 307 |
+
val_dataset = EnhancedTrainingDataset(dataset, val_groups, records_per_group=16)
|
| 308 |
+
|
| 309 |
+
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
|
| 310 |
+
shuffle=True, num_workers=0, collate_fn=collate_fn)
|
| 311 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
|
| 312 |
+
shuffle=False, num_workers=0, collate_fn=collate_fn)
|
| 313 |
+
|
| 314 |
+
# Get dimensions
|
| 315 |
+
sample = train_dataset[0]
|
| 316 |
+
obs_dim = sample["observation"].shape[0]
|
| 317 |
+
action_dim = sample["actions"].shape[1]
|
| 318 |
+
|
| 319 |
+
print(f"Observation dim: {obs_dim}, Action dim: {action_dim}")
|
| 320 |
+
print()
|
| 321 |
+
|
| 322 |
+
# Create enhanced model
|
| 323 |
+
model = DoVLAAttentionEnhanced(
|
| 324 |
+
obs_dim=obs_dim,
|
| 325 |
+
action_dim=action_dim,
|
| 326 |
+
hidden_dim=args.hidden_dim,
|
| 327 |
+
n_heads=args.n_heads,
|
| 328 |
+
n_layers=args.n_layers,
|
| 329 |
+
num_tasks=6,
|
| 330 |
+
use_contrastive=True,
|
| 331 |
+
use_graph=True,
|
| 332 |
+
use_task_adaptive=True
|
| 333 |
+
).to(device)
|
| 334 |
+
|
| 335 |
+
num_params = sum(p.numel() for p in model.parameters())
|
| 336 |
+
print(f"Model parameters: {num_params:,}")
|
| 337 |
+
print()
|
| 338 |
+
|
| 339 |
+
# Optimizer
|
| 340 |
+
optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
|
| 341 |
+
|
| 342 |
+
# Learning rate scheduler
|
| 343 |
+
scheduler = optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=args.epochs)
|
| 344 |
+
|
| 345 |
+
# Training
|
| 346 |
+
best_acc = 0.0
|
| 347 |
+
history = []
|
| 348 |
+
|
| 349 |
+
print("Starting training...")
|
| 350 |
+
print()
|
| 351 |
+
|
| 352 |
+
for epoch in range(args.epochs):
|
| 353 |
+
train_metrics = train_epoch(model, train_loader, optimizer, device, args.contrastive_weight)
|
| 354 |
+
val_metrics = evaluate(model, val_loader, device)
|
| 355 |
+
scheduler.step()
|
| 356 |
+
|
| 357 |
+
val_acc = val_metrics["accuracy"]
|
| 358 |
+
|
| 359 |
+
history.append({
|
| 360 |
+
"epoch": epoch + 1,
|
| 361 |
+
"ranking_loss": train_metrics["ranking_loss"],
|
| 362 |
+
"contrastive_loss": train_metrics["contrastive_loss"],
|
| 363 |
+
"val_accuracy": val_acc,
|
| 364 |
+
"lr": scheduler.get_last_lr()[0]
|
| 365 |
+
})
|
| 366 |
+
|
| 367 |
+
print(f"Epoch {epoch+1:3d}/{args.epochs}: "
|
| 368 |
+
f"rank_loss={train_metrics['ranking_loss']:.4f}, "
|
| 369 |
+
f"contr_loss={train_metrics['contrastive_loss']:.4f}, "
|
| 370 |
+
f"val_acc={val_acc:.4f}")
|
| 371 |
+
|
| 372 |
+
if val_acc > best_acc:
|
| 373 |
+
best_acc = val_acc
|
| 374 |
+
torch.save({
|
| 375 |
+
"model_state_dict": model.state_dict(),
|
| 376 |
+
"epoch": epoch + 1,
|
| 377 |
+
"val_accuracy": val_acc,
|
| 378 |
+
"args": vars(args)
|
| 379 |
+
}, args.out / "best.pt")
|
| 380 |
+
|
| 381 |
+
print()
|
| 382 |
+
print(f"✅ Training complete! Best val accuracy: {best_acc:.4f}")
|
| 383 |
+
|
| 384 |
+
# Save
|
| 385 |
+
with open(args.out / "history.json", "w") as f:
|
| 386 |
+
json.dump(history, f, indent=2)
|
| 387 |
+
|
| 388 |
+
with open(args.out / "config.json", "w") as f:
|
| 389 |
+
json.dump({
|
| 390 |
+
"model": "DoVLA-Attention-Enhanced",
|
| 391 |
+
"components": ["hierarchical_attention", "graph_nn", "contrastive", "task_adaptive"],
|
| 392 |
+
"architecture": {
|
| 393 |
+
"hidden_dim": args.hidden_dim,
|
| 394 |
+
"n_heads": args.n_heads,
|
| 395 |
+
"n_layers": args.n_layers
|
| 396 |
+
},
|
| 397 |
+
"results": {
|
| 398 |
+
"best_val_accuracy": best_acc,
|
| 399 |
+
"num_parameters": num_params
|
| 400 |
+
}
|
| 401 |
+
}, f, indent=2)
|
| 402 |
+
|
| 403 |
+
return 0
|
| 404 |
+
|
| 405 |
+
|
| 406 |
+
if __name__ == "__main__":
|
| 407 |
+
sys.exit(main())
|
scripts/train_dovla_transformer.py
ADDED
|
@@ -0,0 +1,366 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""
|
| 3 |
+
Train DoVLA-Transformer with proper Transformer training recipe.
|
| 4 |
+
|
| 5 |
+
Key improvements over failed Enhanced:
|
| 6 |
+
1. Higher learning rate (0.001 vs 0.0003)
|
| 7 |
+
2. Warmup scheduler (standard for Transformer)
|
| 8 |
+
3. Gradient clipping: 1.0 (standard)
|
| 9 |
+
4. Pure ranking loss (no contrastive)
|
| 10 |
+
5. Standard Transformer components (proven)
|
| 11 |
+
|
| 12 |
+
Expected: 42-47% success
|
| 13 |
+
"""
|
| 14 |
+
from __future__ import annotations
|
| 15 |
+
|
| 16 |
+
import argparse
|
| 17 |
+
import json
|
| 18 |
+
import random
|
| 19 |
+
import sys
|
| 20 |
+
from pathlib import Path
|
| 21 |
+
|
| 22 |
+
import numpy as np
|
| 23 |
+
import torch
|
| 24 |
+
import torch.nn as nn
|
| 25 |
+
import torch.optim as optim
|
| 26 |
+
from torch.utils.data import DataLoader, Dataset
|
| 27 |
+
|
| 28 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 29 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 30 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 31 |
+
|
| 32 |
+
from dovla_cil.models.dovla_transformer import DoVLATransformer
|
| 33 |
+
from dovla_cil.data.datasets import CILDataset
|
| 34 |
+
|
| 35 |
+
|
| 36 |
+
class TransformerTrainingDataset(Dataset):
|
| 37 |
+
"""Dataset for DoVLA-Transformer training."""
|
| 38 |
+
|
| 39 |
+
def __init__(self, dataset: CILDataset, group_ids: list[str],
|
| 40 |
+
records_per_group: int = 16, max_obs_dim: int = 70, max_act_dim: int = 32):
|
| 41 |
+
self.dataset = dataset
|
| 42 |
+
self.group_ids = group_ids
|
| 43 |
+
self.records_per_group = records_per_group
|
| 44 |
+
self.max_obs_dim = max_obs_dim
|
| 45 |
+
self.max_act_dim = max_act_dim
|
| 46 |
+
|
| 47 |
+
def _pad(self, vec: list[float], target: int) -> list[float]:
|
| 48 |
+
if len(vec) >= target:
|
| 49 |
+
return vec[:target]
|
| 50 |
+
return vec + [0.0] * (target - len(vec))
|
| 51 |
+
|
| 52 |
+
def __len__(self):
|
| 53 |
+
return len(self.group_ids)
|
| 54 |
+
|
| 55 |
+
def __getitem__(self, idx):
|
| 56 |
+
group_id = self.group_ids[idx]
|
| 57 |
+
records = self.dataset.get_group(group_id)
|
| 58 |
+
|
| 59 |
+
if len(records) > self.records_per_group:
|
| 60 |
+
records = random.sample(records, self.records_per_group)
|
| 61 |
+
|
| 62 |
+
# Observation
|
| 63 |
+
obs_data = records[0].observation_inline
|
| 64 |
+
if "state" in obs_data:
|
| 65 |
+
obs = list(obs_data["state"])
|
| 66 |
+
else:
|
| 67 |
+
obs = []
|
| 68 |
+
for v in obs_data.values():
|
| 69 |
+
if isinstance(v, list):
|
| 70 |
+
obs.extend(v)
|
| 71 |
+
elif isinstance(v, (int, float)):
|
| 72 |
+
obs.append(v)
|
| 73 |
+
obs = self._pad([float(x) for x in obs], self.max_obs_dim)
|
| 74 |
+
|
| 75 |
+
# Actions
|
| 76 |
+
actions = [self._pad(r.action_chunk.flat_values, self.max_act_dim) for r in records]
|
| 77 |
+
|
| 78 |
+
# Rewards
|
| 79 |
+
rewards = [r.reward.score for r in records]
|
| 80 |
+
|
| 81 |
+
return {
|
| 82 |
+
"observation": torch.FloatTensor(obs),
|
| 83 |
+
"actions": torch.FloatTensor(actions),
|
| 84 |
+
"rewards": torch.FloatTensor(rewards)
|
| 85 |
+
}
|
| 86 |
+
|
| 87 |
+
|
| 88 |
+
def collate_fn(batch):
|
| 89 |
+
"""Collate with padding to max K in batch."""
|
| 90 |
+
max_k = max(b["actions"].shape[0] for b in batch)
|
| 91 |
+
batch_size = len(batch)
|
| 92 |
+
obs_dim = batch[0]["observation"].shape[0]
|
| 93 |
+
action_dim = batch[0]["actions"].shape[1]
|
| 94 |
+
|
| 95 |
+
obs_batch = torch.stack([b["observation"] for b in batch])
|
| 96 |
+
actions_batch = torch.zeros(batch_size, max_k, action_dim)
|
| 97 |
+
rewards_batch = torch.zeros(batch_size, max_k)
|
| 98 |
+
|
| 99 |
+
for i, b in enumerate(batch):
|
| 100 |
+
k = b["actions"].shape[0]
|
| 101 |
+
actions_batch[i, :k] = b["actions"]
|
| 102 |
+
rewards_batch[i, :k] = b["rewards"]
|
| 103 |
+
|
| 104 |
+
return {
|
| 105 |
+
"observation": obs_batch,
|
| 106 |
+
"actions": actions_batch,
|
| 107 |
+
"rewards": rewards_batch
|
| 108 |
+
}
|
| 109 |
+
|
| 110 |
+
|
| 111 |
+
def set_seed(seed: int):
|
| 112 |
+
random.seed(seed)
|
| 113 |
+
np.random.seed(seed)
|
| 114 |
+
torch.manual_seed(seed)
|
| 115 |
+
if torch.cuda.is_available():
|
| 116 |
+
torch.cuda.manual_seed_all(seed)
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps):
|
| 120 |
+
"""Cosine schedule with linear warmup (standard for Transformer)."""
|
| 121 |
+
def lr_lambda(current_step):
|
| 122 |
+
if current_step < num_warmup_steps:
|
| 123 |
+
return float(current_step) / float(max(1, num_warmup_steps))
|
| 124 |
+
progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))
|
| 125 |
+
return max(0.0, 0.5 * (1.0 + np.cos(np.pi * progress)))
|
| 126 |
+
|
| 127 |
+
return optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def train_epoch(model: nn.Module, dataloader: DataLoader,
|
| 131 |
+
optimizer: optim.Optimizer, device: torch.device) -> float:
|
| 132 |
+
"""Train for one epoch."""
|
| 133 |
+
model.train()
|
| 134 |
+
total_loss = 0.0
|
| 135 |
+
num_batches = 0
|
| 136 |
+
|
| 137 |
+
for batch in dataloader:
|
| 138 |
+
obs = batch["observation"].to(device)
|
| 139 |
+
actions = batch["actions"].to(device)
|
| 140 |
+
rewards = batch["rewards"].to(device)
|
| 141 |
+
|
| 142 |
+
optimizer.zero_grad()
|
| 143 |
+
|
| 144 |
+
scores = model(obs, actions)
|
| 145 |
+
|
| 146 |
+
# Ranking loss
|
| 147 |
+
batch_size, K = rewards.shape
|
| 148 |
+
loss = 0.0
|
| 149 |
+
count = 0
|
| 150 |
+
|
| 151 |
+
for b in range(batch_size):
|
| 152 |
+
for i in range(K):
|
| 153 |
+
for j in range(K):
|
| 154 |
+
if i != j and rewards[b, i] != rewards[b, j]:
|
| 155 |
+
target = 1.0 if rewards[b, i] > rewards[b, j] else 0.0
|
| 156 |
+
pred = torch.sigmoid(scores[b, i, j])
|
| 157 |
+
loss += nn.functional.binary_cross_entropy(
|
| 158 |
+
pred.unsqueeze(0),
|
| 159 |
+
torch.tensor([target], device=device)
|
| 160 |
+
)
|
| 161 |
+
count += 1
|
| 162 |
+
|
| 163 |
+
if count > 0:
|
| 164 |
+
loss = loss / count
|
| 165 |
+
loss.backward()
|
| 166 |
+
|
| 167 |
+
# Gradient clipping
|
| 168 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 169 |
+
|
| 170 |
+
optimizer.step()
|
| 171 |
+
|
| 172 |
+
total_loss += loss.item()
|
| 173 |
+
num_batches += 1
|
| 174 |
+
|
| 175 |
+
return total_loss / max(num_batches, 1)
|
| 176 |
+
|
| 177 |
+
|
| 178 |
+
def evaluate(model: nn.Module, dataloader: DataLoader, device: torch.device) -> dict:
|
| 179 |
+
"""Evaluate model - proper action selection accuracy."""
|
| 180 |
+
model.eval()
|
| 181 |
+
correct_top1 = 0
|
| 182 |
+
total_groups = 0
|
| 183 |
+
|
| 184 |
+
with torch.no_grad():
|
| 185 |
+
for batch in dataloader:
|
| 186 |
+
obs = batch["observation"].to(device)
|
| 187 |
+
actions = batch["actions"].to(device)
|
| 188 |
+
rewards = batch["rewards"].to(device)
|
| 189 |
+
|
| 190 |
+
scores_matrix = model(obs, actions)
|
| 191 |
+
|
| 192 |
+
# Aggregate pairwise to per-action scores
|
| 193 |
+
batch_size, K = rewards.shape
|
| 194 |
+
for b in range(batch_size):
|
| 195 |
+
# Sum wins for each action
|
| 196 |
+
action_scores = scores_matrix[b].sum(dim=1).cpu().tolist()
|
| 197 |
+
|
| 198 |
+
# Select best
|
| 199 |
+
selected = max(range(K), key=lambda i: action_scores[i])
|
| 200 |
+
best_utility = max(rewards[b].cpu().tolist())
|
| 201 |
+
|
| 202 |
+
if abs(rewards[b, selected].item() - best_utility) < 1e-6:
|
| 203 |
+
correct_top1 += 1
|
| 204 |
+
total_groups += 1
|
| 205 |
+
|
| 206 |
+
return {
|
| 207 |
+
"top1_accuracy": correct_top1 / max(total_groups, 1)
|
| 208 |
+
}
|
| 209 |
+
|
| 210 |
+
|
| 211 |
+
def main(argv: list[str] | None = None) -> int:
|
| 212 |
+
parser = argparse.ArgumentParser(description="Train DoVLA-Transformer")
|
| 213 |
+
|
| 214 |
+
# Data
|
| 215 |
+
parser.add_argument("--dataset", type=Path, required=True)
|
| 216 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 217 |
+
|
| 218 |
+
# Architecture
|
| 219 |
+
parser.add_argument("--d-model", type=int, default=256)
|
| 220 |
+
parser.add_argument("--n-heads", type=int, default=8)
|
| 221 |
+
parser.add_argument("--n-layers", type=int, default=3)
|
| 222 |
+
parser.add_argument("--d-ff", type=int, default=1024)
|
| 223 |
+
|
| 224 |
+
# Training
|
| 225 |
+
parser.add_argument("--epochs", type=int, default=50)
|
| 226 |
+
parser.add_argument("--batch-size", type=int, default=16)
|
| 227 |
+
parser.add_argument("--lr", type=float, default=0.001) # Higher than failed Enhanced
|
| 228 |
+
parser.add_argument("--weight-decay", type=float, default=0.01)
|
| 229 |
+
parser.add_argument("--warmup-steps", type=int, default=500)
|
| 230 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 231 |
+
parser.add_argument("--val-fraction", type=float, default=0.2)
|
| 232 |
+
|
| 233 |
+
# System
|
| 234 |
+
parser.add_argument("--device", default="auto")
|
| 235 |
+
|
| 236 |
+
args = parser.parse_args(argv)
|
| 237 |
+
|
| 238 |
+
set_seed(args.seed)
|
| 239 |
+
args.out.mkdir(parents=True, exist_ok=True)
|
| 240 |
+
|
| 241 |
+
if args.device == "auto":
|
| 242 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 243 |
+
else:
|
| 244 |
+
device = torch.device(args.device)
|
| 245 |
+
|
| 246 |
+
print("=" * 70)
|
| 247 |
+
print("DoVLA-Transformer Training (BREAKTHROUGH)")
|
| 248 |
+
print("=" * 70)
|
| 249 |
+
print(f"Dataset: {args.dataset}")
|
| 250 |
+
print(f"Device: {device}")
|
| 251 |
+
print(f"Architecture: Pure Transformer (d={args.d_model}, heads={args.n_heads}, layers={args.n_layers})")
|
| 252 |
+
print(f"LR: {args.lr} (higher than failed Enhanced 0.0003)")
|
| 253 |
+
print(f"Warmup: {args.warmup_steps} steps")
|
| 254 |
+
print(f"Seed: {args.seed}")
|
| 255 |
+
print()
|
| 256 |
+
|
| 257 |
+
# Load data
|
| 258 |
+
print("Loading dataset...")
|
| 259 |
+
dataset = CILDataset(args.dataset)
|
| 260 |
+
all_groups = list(dataset.group_ids)
|
| 261 |
+
|
| 262 |
+
random.shuffle(all_groups)
|
| 263 |
+
split_idx = int(len(all_groups) * (1 - args.val_fraction))
|
| 264 |
+
train_groups = all_groups[:split_idx]
|
| 265 |
+
val_groups = all_groups[split_idx:]
|
| 266 |
+
|
| 267 |
+
print(f"Total: {len(all_groups)}, Train: {len(train_groups)}, Val: {len(val_groups)}")
|
| 268 |
+
print()
|
| 269 |
+
|
| 270 |
+
# Datasets
|
| 271 |
+
train_dataset = TransformerTrainingDataset(dataset, train_groups)
|
| 272 |
+
val_dataset = TransformerTrainingDataset(dataset, val_groups)
|
| 273 |
+
|
| 274 |
+
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
|
| 275 |
+
shuffle=True, num_workers=0, collate_fn=collate_fn)
|
| 276 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
|
| 277 |
+
shuffle=False, num_workers=0, collate_fn=collate_fn)
|
| 278 |
+
|
| 279 |
+
# Model
|
| 280 |
+
model = DoVLATransformer(
|
| 281 |
+
obs_dim=70,
|
| 282 |
+
action_dim=32,
|
| 283 |
+
lang_dim=0,
|
| 284 |
+
d_model=args.d_model,
|
| 285 |
+
n_heads=args.n_heads,
|
| 286 |
+
n_layers=args.n_layers,
|
| 287 |
+
d_ff=args.d_ff,
|
| 288 |
+
dropout=0.1
|
| 289 |
+
).to(device)
|
| 290 |
+
|
| 291 |
+
num_params = sum(p.numel() for p in model.parameters())
|
| 292 |
+
print(f"Model parameters: {num_params:,}")
|
| 293 |
+
print()
|
| 294 |
+
|
| 295 |
+
# Optimizer & Scheduler
|
| 296 |
+
optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
|
| 297 |
+
|
| 298 |
+
num_training_steps = len(train_loader) * args.epochs
|
| 299 |
+
scheduler = get_cosine_schedule_with_warmup(optimizer, args.warmup_steps, num_training_steps)
|
| 300 |
+
|
| 301 |
+
# Training
|
| 302 |
+
best_acc = 0.0
|
| 303 |
+
history = []
|
| 304 |
+
|
| 305 |
+
print("Starting training...")
|
| 306 |
+
print()
|
| 307 |
+
|
| 308 |
+
for epoch in range(args.epochs):
|
| 309 |
+
train_loss = train_epoch(model, train_loader, optimizer, device)
|
| 310 |
+
val_metrics = evaluate(model, val_loader, device)
|
| 311 |
+
scheduler.step()
|
| 312 |
+
|
| 313 |
+
val_acc = val_metrics["top1_accuracy"]
|
| 314 |
+
|
| 315 |
+
history.append({
|
| 316 |
+
"epoch": epoch + 1,
|
| 317 |
+
"train_loss": train_loss,
|
| 318 |
+
"val_top1_accuracy": val_acc,
|
| 319 |
+
"lr": scheduler.get_last_lr()[0]
|
| 320 |
+
})
|
| 321 |
+
|
| 322 |
+
print(f"Epoch {epoch+1:3d}/{args.epochs}: "
|
| 323 |
+
f"loss={train_loss:.4f}, val_top1={val_acc:.4f}, lr={scheduler.get_last_lr()[0]:.6f}")
|
| 324 |
+
|
| 325 |
+
if val_acc > best_acc:
|
| 326 |
+
best_acc = val_acc
|
| 327 |
+
torch.save({
|
| 328 |
+
"model_state_dict": model.state_dict(),
|
| 329 |
+
"epoch": epoch + 1,
|
| 330 |
+
"val_top1_accuracy": val_acc,
|
| 331 |
+
"args": vars(args)
|
| 332 |
+
}, args.out / "best.pt")
|
| 333 |
+
|
| 334 |
+
print()
|
| 335 |
+
print(f"✅ Training complete! Best val top-1: {best_acc:.4f}")
|
| 336 |
+
|
| 337 |
+
# Save
|
| 338 |
+
with open(args.out / "history.json", "w") as f:
|
| 339 |
+
json.dump(history, f, indent=2)
|
| 340 |
+
|
| 341 |
+
with open(args.out / "config.json", "w") as f:
|
| 342 |
+
json.dump({
|
| 343 |
+
"model": "DoVLA-Transformer",
|
| 344 |
+
"architecture": {
|
| 345 |
+
"d_model": args.d_model,
|
| 346 |
+
"n_heads": args.n_heads,
|
| 347 |
+
"n_layers": args.n_layers,
|
| 348 |
+
"d_ff": args.d_ff
|
| 349 |
+
},
|
| 350 |
+
"training": {
|
| 351 |
+
"lr": args.lr,
|
| 352 |
+
"warmup_steps": args.warmup_steps,
|
| 353 |
+
"epochs": args.epochs,
|
| 354 |
+
"seed": args.seed
|
| 355 |
+
},
|
| 356 |
+
"results": {
|
| 357 |
+
"best_val_top1_accuracy": best_acc,
|
| 358 |
+
"num_parameters": num_params
|
| 359 |
+
}
|
| 360 |
+
}, f, indent=2)
|
| 361 |
+
|
| 362 |
+
return 0
|
| 363 |
+
|
| 364 |
+
|
| 365 |
+
if __name__ == "__main__":
|
| 366 |
+
sys.exit(main())
|
scripts/train_hybrid_direct.py
ADDED
|
@@ -0,0 +1,348 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""
|
| 3 |
+
Train DoVLA-Hybrid with DIRECT scoring (NOT pairwise).
|
| 4 |
+
|
| 5 |
+
Key improvement: Predict reward + success directly
|
| 6 |
+
Expected: 45-48% baseline (vs 37% pairwise)
|
| 7 |
+
"""
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import argparse
|
| 11 |
+
import json
|
| 12 |
+
import random
|
| 13 |
+
import sys
|
| 14 |
+
from pathlib import Path
|
| 15 |
+
|
| 16 |
+
import numpy as np
|
| 17 |
+
import torch
|
| 18 |
+
import torch.nn as nn
|
| 19 |
+
import torch.nn.functional as F
|
| 20 |
+
import torch.optim as optim
|
| 21 |
+
from torch.utils.data import DataLoader, Dataset
|
| 22 |
+
|
| 23 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 24 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 25 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 26 |
+
|
| 27 |
+
from dovla_cil.models.dovla_hybrid import DoVLAHybrid
|
| 28 |
+
from dovla_cil.data.datasets import CILDataset
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class HybridTrainingDataset(Dataset):
|
| 32 |
+
"""Dataset for hybrid direct scoring."""
|
| 33 |
+
|
| 34 |
+
def __init__(self, dataset: CILDataset, group_ids: list[str],
|
| 35 |
+
records_per_group: int = 16, max_obs_dim: int = 70, max_act_dim: int = 32):
|
| 36 |
+
self.dataset = dataset
|
| 37 |
+
self.group_ids = group_ids
|
| 38 |
+
self.records_per_group = records_per_group
|
| 39 |
+
self.max_obs_dim = max_obs_dim
|
| 40 |
+
self.max_act_dim = max_act_dim
|
| 41 |
+
|
| 42 |
+
def _pad(self, vec: list[float], target: int) -> list[float]:
|
| 43 |
+
if len(vec) >= target:
|
| 44 |
+
return vec[:target]
|
| 45 |
+
return vec + [0.0] * (target - len(vec))
|
| 46 |
+
|
| 47 |
+
def __len__(self):
|
| 48 |
+
return len(self.group_ids)
|
| 49 |
+
|
| 50 |
+
def __getitem__(self, idx):
|
| 51 |
+
group_id = self.group_ids[idx]
|
| 52 |
+
records = self.dataset.get_group(group_id)
|
| 53 |
+
|
| 54 |
+
if len(records) > self.records_per_group:
|
| 55 |
+
records = random.sample(records, self.records_per_group)
|
| 56 |
+
|
| 57 |
+
# Observation
|
| 58 |
+
obs_data = records[0].observation_inline
|
| 59 |
+
if "state" in obs_data:
|
| 60 |
+
obs = list(obs_data["state"])
|
| 61 |
+
else:
|
| 62 |
+
obs = []
|
| 63 |
+
for v in obs_data.values():
|
| 64 |
+
if isinstance(v, list):
|
| 65 |
+
obs.extend(v)
|
| 66 |
+
elif isinstance(v, (int, float)):
|
| 67 |
+
obs.append(v)
|
| 68 |
+
obs = self._pad([float(x) for x in obs], self.max_obs_dim)
|
| 69 |
+
|
| 70 |
+
# Actions
|
| 71 |
+
actions = [self._pad(r.action_chunk.flat_values, self.max_act_dim) for r in records]
|
| 72 |
+
|
| 73 |
+
# Rewards (direct targets!)
|
| 74 |
+
rewards = [r.reward.score for r in records]
|
| 75 |
+
|
| 76 |
+
# Success labels (direct targets!)
|
| 77 |
+
successes = [float(r.reward.terminal_success) for r in records]
|
| 78 |
+
|
| 79 |
+
return {
|
| 80 |
+
"observation": torch.FloatTensor(obs),
|
| 81 |
+
"actions": torch.FloatTensor(actions),
|
| 82 |
+
"rewards": torch.FloatTensor(rewards),
|
| 83 |
+
"successes": torch.FloatTensor(successes)
|
| 84 |
+
}
|
| 85 |
+
|
| 86 |
+
|
| 87 |
+
def collate_fn(batch):
|
| 88 |
+
"""Collate with padding."""
|
| 89 |
+
max_k = max(b["actions"].shape[0] for b in batch)
|
| 90 |
+
batch_size = len(batch)
|
| 91 |
+
obs_dim = batch[0]["observation"].shape[0]
|
| 92 |
+
action_dim = batch[0]["actions"].shape[1]
|
| 93 |
+
|
| 94 |
+
obs_batch = torch.stack([b["observation"] for b in batch])
|
| 95 |
+
actions_batch = torch.zeros(batch_size, max_k, action_dim)
|
| 96 |
+
rewards_batch = torch.zeros(batch_size, max_k)
|
| 97 |
+
successes_batch = torch.zeros(batch_size, max_k)
|
| 98 |
+
|
| 99 |
+
for i, b in enumerate(batch):
|
| 100 |
+
k = b["actions"].shape[0]
|
| 101 |
+
actions_batch[i, :k] = b["actions"]
|
| 102 |
+
rewards_batch[i, :k] = b["rewards"]
|
| 103 |
+
successes_batch[i, :k] = b["successes"]
|
| 104 |
+
|
| 105 |
+
return {
|
| 106 |
+
"observation": obs_batch,
|
| 107 |
+
"actions": actions_batch,
|
| 108 |
+
"rewards": rewards_batch,
|
| 109 |
+
"successes": successes_batch
|
| 110 |
+
}
|
| 111 |
+
|
| 112 |
+
|
| 113 |
+
def set_seed(seed: int):
|
| 114 |
+
random.seed(seed)
|
| 115 |
+
np.random.seed(seed)
|
| 116 |
+
torch.manual_seed(seed)
|
| 117 |
+
if torch.cuda.is_available():
|
| 118 |
+
torch.cuda.manual_seed_all(seed)
|
| 119 |
+
|
| 120 |
+
|
| 121 |
+
def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps):
|
| 122 |
+
def lr_lambda(current_step):
|
| 123 |
+
if current_step < num_warmup_steps:
|
| 124 |
+
return float(current_step) / float(max(1, num_warmup_steps))
|
| 125 |
+
progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))
|
| 126 |
+
return max(0.0, 0.5 * (1.0 + np.cos(np.pi * progress)))
|
| 127 |
+
return optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
|
| 128 |
+
|
| 129 |
+
|
| 130 |
+
def train_epoch(model: nn.Module, dataloader: DataLoader,
|
| 131 |
+
optimizer: optim.Optimizer, device: torch.device) -> dict:
|
| 132 |
+
"""Train for one epoch with DIRECT scoring."""
|
| 133 |
+
model.train()
|
| 134 |
+
total_reward_loss = 0.0
|
| 135 |
+
total_success_loss = 0.0
|
| 136 |
+
num_batches = 0
|
| 137 |
+
|
| 138 |
+
for batch in dataloader:
|
| 139 |
+
obs = batch["observation"].to(device)
|
| 140 |
+
actions = batch["actions"].to(device)
|
| 141 |
+
target_rewards = batch["rewards"].to(device)
|
| 142 |
+
target_successes = batch["successes"].to(device)
|
| 143 |
+
|
| 144 |
+
optimizer.zero_grad()
|
| 145 |
+
|
| 146 |
+
# Forward: predict rewards and success probs DIRECTLY
|
| 147 |
+
pred_rewards, pred_success_probs = model(obs, actions)
|
| 148 |
+
|
| 149 |
+
# Loss 1: MSE for reward prediction
|
| 150 |
+
reward_loss = F.mse_loss(pred_rewards, target_rewards)
|
| 151 |
+
|
| 152 |
+
# Loss 2: BCE for success prediction
|
| 153 |
+
success_loss = F.binary_cross_entropy(pred_success_probs, target_successes)
|
| 154 |
+
|
| 155 |
+
# Combined loss
|
| 156 |
+
loss = reward_loss + success_loss
|
| 157 |
+
|
| 158 |
+
loss.backward()
|
| 159 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 160 |
+
optimizer.step()
|
| 161 |
+
|
| 162 |
+
total_reward_loss += reward_loss.item()
|
| 163 |
+
total_success_loss += success_loss.item()
|
| 164 |
+
num_batches += 1
|
| 165 |
+
|
| 166 |
+
return {
|
| 167 |
+
"reward_loss": total_reward_loss / max(num_batches, 1),
|
| 168 |
+
"success_loss": total_success_loss / max(num_batches, 1),
|
| 169 |
+
"total_loss": (total_reward_loss + total_success_loss) / max(num_batches, 1)
|
| 170 |
+
}
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
def evaluate(model: nn.Module, dataloader: DataLoader, device: torch.device) -> dict:
|
| 174 |
+
"""Evaluate with DIRECT selection."""
|
| 175 |
+
model.eval()
|
| 176 |
+
correct_top1 = 0
|
| 177 |
+
total_groups = 0
|
| 178 |
+
|
| 179 |
+
with torch.no_grad():
|
| 180 |
+
for batch in dataloader:
|
| 181 |
+
obs = batch["observation"].to(device)
|
| 182 |
+
actions = batch["actions"].to(device)
|
| 183 |
+
target_rewards = batch["rewards"].to(device)
|
| 184 |
+
|
| 185 |
+
# Predict directly
|
| 186 |
+
pred_rewards, pred_success_probs = model(obs, actions)
|
| 187 |
+
|
| 188 |
+
# Hybrid selection: success_prob * predicted_reward
|
| 189 |
+
hybrid_scores = pred_success_probs * pred_rewards
|
| 190 |
+
|
| 191 |
+
batch_size, K = target_rewards.shape
|
| 192 |
+
for b in range(batch_size):
|
| 193 |
+
# Select action with highest hybrid score
|
| 194 |
+
selected = hybrid_scores[b].argmax().item()
|
| 195 |
+
|
| 196 |
+
# Check if selected best reward
|
| 197 |
+
best_reward = target_rewards[b].max().item()
|
| 198 |
+
if abs(target_rewards[b, selected].item() - best_reward) < 1e-6:
|
| 199 |
+
correct_top1 += 1
|
| 200 |
+
total_groups += 1
|
| 201 |
+
|
| 202 |
+
return {"top1_accuracy": correct_top1 / max(total_groups, 1)}
|
| 203 |
+
|
| 204 |
+
|
| 205 |
+
def main(argv: list[str] | None = None) -> int:
|
| 206 |
+
parser = argparse.ArgumentParser(description="Train DoVLA-Hybrid (Direct Scoring)")
|
| 207 |
+
|
| 208 |
+
parser.add_argument("--dataset", type=Path, required=True)
|
| 209 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 210 |
+
parser.add_argument("--d-model", type=int, default=256)
|
| 211 |
+
parser.add_argument("--n-heads", type=int, default=8)
|
| 212 |
+
parser.add_argument("--n-layers", type=int, default=3)
|
| 213 |
+
parser.add_argument("--d-ff", type=int, default=1024)
|
| 214 |
+
parser.add_argument("--epochs", type=int, default=50)
|
| 215 |
+
parser.add_argument("--batch-size", type=int, default=16)
|
| 216 |
+
parser.add_argument("--lr", type=float, default=0.001)
|
| 217 |
+
parser.add_argument("--weight-decay", type=float, default=0.01)
|
| 218 |
+
parser.add_argument("--warmup-steps", type=int, default=500)
|
| 219 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 220 |
+
parser.add_argument("--val-fraction", type=float, default=0.2)
|
| 221 |
+
parser.add_argument("--device", default="auto")
|
| 222 |
+
|
| 223 |
+
args = parser.parse_args(argv)
|
| 224 |
+
set_seed(args.seed)
|
| 225 |
+
args.out.mkdir(parents=True, exist_ok=True)
|
| 226 |
+
|
| 227 |
+
if args.device == "auto":
|
| 228 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 229 |
+
else:
|
| 230 |
+
device = torch.device(args.device)
|
| 231 |
+
|
| 232 |
+
print("=" * 70)
|
| 233 |
+
print("DoVLA-Hybrid: DIRECT Scoring (NOT Pairwise)")
|
| 234 |
+
print("=" * 70)
|
| 235 |
+
print(f"Dataset: {args.dataset}")
|
| 236 |
+
print(f"Device: {device}")
|
| 237 |
+
print(f"Approach: Predict reward + success DIRECTLY")
|
| 238 |
+
print(f"Expected: 45-48% (vs 37% pairwise baseline)")
|
| 239 |
+
print()
|
| 240 |
+
|
| 241 |
+
# Load data
|
| 242 |
+
dataset = CILDataset(args.dataset)
|
| 243 |
+
all_groups = list(dataset.group_ids)
|
| 244 |
+
random.shuffle(all_groups)
|
| 245 |
+
split_idx = int(len(all_groups) * (1 - args.val_fraction))
|
| 246 |
+
train_groups = all_groups[:split_idx]
|
| 247 |
+
val_groups = all_groups[split_idx:]
|
| 248 |
+
|
| 249 |
+
print(f"Total: {len(all_groups)}, Train: {len(train_groups)}, Val: {len(val_groups)}")
|
| 250 |
+
print()
|
| 251 |
+
|
| 252 |
+
train_dataset = HybridTrainingDataset(dataset, train_groups)
|
| 253 |
+
val_dataset = HybridTrainingDataset(dataset, val_groups)
|
| 254 |
+
|
| 255 |
+
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
|
| 256 |
+
shuffle=True, num_workers=0, collate_fn=collate_fn)
|
| 257 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
|
| 258 |
+
shuffle=False, num_workers=0, collate_fn=collate_fn)
|
| 259 |
+
|
| 260 |
+
# Model
|
| 261 |
+
model = DoVLAHybrid(
|
| 262 |
+
obs_dim=70,
|
| 263 |
+
action_dim=32,
|
| 264 |
+
lang_dim=0,
|
| 265 |
+
d_model=args.d_model,
|
| 266 |
+
n_heads=args.n_heads,
|
| 267 |
+
n_layers=args.n_layers,
|
| 268 |
+
d_ff=args.d_ff,
|
| 269 |
+
dropout=0.1
|
| 270 |
+
).to(device)
|
| 271 |
+
|
| 272 |
+
num_params = sum(p.numel() for p in model.parameters())
|
| 273 |
+
print(f"Model parameters: {num_params:,}")
|
| 274 |
+
print()
|
| 275 |
+
|
| 276 |
+
optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
|
| 277 |
+
num_training_steps = len(train_loader) * args.epochs
|
| 278 |
+
scheduler = get_cosine_schedule_with_warmup(optimizer, args.warmup_steps, num_training_steps)
|
| 279 |
+
|
| 280 |
+
best_acc = 0.0
|
| 281 |
+
history = []
|
| 282 |
+
|
| 283 |
+
print("Starting training...")
|
| 284 |
+
print()
|
| 285 |
+
|
| 286 |
+
for epoch in range(args.epochs):
|
| 287 |
+
train_metrics = train_epoch(model, train_loader, optimizer, device)
|
| 288 |
+
val_metrics = evaluate(model, val_loader, device)
|
| 289 |
+
scheduler.step()
|
| 290 |
+
|
| 291 |
+
val_acc = val_metrics["top1_accuracy"]
|
| 292 |
+
|
| 293 |
+
history.append({
|
| 294 |
+
"epoch": epoch + 1,
|
| 295 |
+
"train_reward_loss": train_metrics["reward_loss"],
|
| 296 |
+
"train_success_loss": train_metrics["success_loss"],
|
| 297 |
+
"train_total_loss": train_metrics["total_loss"],
|
| 298 |
+
"val_top1_accuracy": val_acc,
|
| 299 |
+
"lr": scheduler.get_last_lr()[0]
|
| 300 |
+
})
|
| 301 |
+
|
| 302 |
+
print(f"Epoch {epoch+1:3d}/{args.epochs}: "
|
| 303 |
+
f"r_loss={train_metrics['reward_loss']:.4f}, "
|
| 304 |
+
f"s_loss={train_metrics['success_loss']:.4f}, "
|
| 305 |
+
f"val_top1={val_acc:.4f}")
|
| 306 |
+
|
| 307 |
+
if val_acc > best_acc:
|
| 308 |
+
best_acc = val_acc
|
| 309 |
+
torch.save({
|
| 310 |
+
"model_state_dict": model.state_dict(),
|
| 311 |
+
"epoch": epoch + 1,
|
| 312 |
+
"val_top1_accuracy": val_acc,
|
| 313 |
+
"args": vars(args)
|
| 314 |
+
}, args.out / "best.pt")
|
| 315 |
+
|
| 316 |
+
print()
|
| 317 |
+
print(f"✅ Training complete! Best val top-1: {best_acc:.4f}")
|
| 318 |
+
|
| 319 |
+
with open(args.out / "history.json", "w") as f:
|
| 320 |
+
json.dump(history, f, indent=2)
|
| 321 |
+
|
| 322 |
+
with open(args.out / "config.json", "w") as f:
|
| 323 |
+
json.dump({
|
| 324 |
+
"model": "DoVLA-Hybrid-Direct",
|
| 325 |
+
"approach": "direct_scoring",
|
| 326 |
+
"architecture": {
|
| 327 |
+
"d_model": args.d_model,
|
| 328 |
+
"n_heads": args.n_heads,
|
| 329 |
+
"n_layers": args.n_layers,
|
| 330 |
+
"d_ff": args.d_ff
|
| 331 |
+
},
|
| 332 |
+
"training": {
|
| 333 |
+
"lr": args.lr,
|
| 334 |
+
"warmup_steps": args.warmup_steps,
|
| 335 |
+
"epochs": args.epochs,
|
| 336 |
+
"seed": args.seed
|
| 337 |
+
},
|
| 338 |
+
"results": {
|
| 339 |
+
"best_val_top1_accuracy": best_acc,
|
| 340 |
+
"num_parameters": num_params
|
| 341 |
+
}
|
| 342 |
+
}, f, indent=2)
|
| 343 |
+
|
| 344 |
+
return 0
|
| 345 |
+
|
| 346 |
+
|
| 347 |
+
if __name__ == "__main__":
|
| 348 |
+
sys.exit(main())
|
scripts/train_transformer_with_language.py
ADDED
|
@@ -0,0 +1,368 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
#!/usr/bin/env python
|
| 2 |
+
"""
|
| 3 |
+
Train DoVLA-Transformer WITH LANGUAGE EMBEDDINGS.
|
| 4 |
+
|
| 5 |
+
This is the improved version that uses instruction embeddings.
|
| 6 |
+
Expected improvement: +8-11% (50-55% from 42-44% baseline)
|
| 7 |
+
"""
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
import argparse
|
| 11 |
+
import json
|
| 12 |
+
import pickle
|
| 13 |
+
import random
|
| 14 |
+
import sys
|
| 15 |
+
from pathlib import Path
|
| 16 |
+
|
| 17 |
+
import numpy as np
|
| 18 |
+
import torch
|
| 19 |
+
import torch.nn as nn
|
| 20 |
+
import torch.optim as optim
|
| 21 |
+
from torch.utils.data import DataLoader, Dataset
|
| 22 |
+
|
| 23 |
+
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
| 24 |
+
if str(PROJECT_ROOT) not in sys.path:
|
| 25 |
+
sys.path.insert(0, str(PROJECT_ROOT))
|
| 26 |
+
|
| 27 |
+
from dovla_cil.models.dovla_transformer import DoVLATransformer
|
| 28 |
+
from dovla_cil.data.datasets import CILDataset
|
| 29 |
+
|
| 30 |
+
|
| 31 |
+
class TransformerLangDataset(Dataset):
|
| 32 |
+
"""Dataset for DoVLA-Transformer with language embeddings."""
|
| 33 |
+
|
| 34 |
+
def __init__(self, dataset: CILDataset, group_ids: list[str],
|
| 35 |
+
embeddings: dict, records_per_group: int = 16,
|
| 36 |
+
max_obs_dim: int = 70, max_act_dim: int = 32):
|
| 37 |
+
self.dataset = dataset
|
| 38 |
+
self.group_ids = group_ids
|
| 39 |
+
self.embeddings = embeddings # {group_id: embedding}
|
| 40 |
+
self.records_per_group = records_per_group
|
| 41 |
+
self.max_obs_dim = max_obs_dim
|
| 42 |
+
self.max_act_dim = max_act_dim
|
| 43 |
+
|
| 44 |
+
def _pad(self, vec: list[float], target: int) -> list[float]:
|
| 45 |
+
if len(vec) >= target:
|
| 46 |
+
return vec[:target]
|
| 47 |
+
return vec + [0.0] * (target - len(vec))
|
| 48 |
+
|
| 49 |
+
def __len__(self):
|
| 50 |
+
return len(self.group_ids)
|
| 51 |
+
|
| 52 |
+
def __getitem__(self, idx):
|
| 53 |
+
group_id = self.group_ids[idx]
|
| 54 |
+
records = self.dataset.get_group(group_id)
|
| 55 |
+
|
| 56 |
+
if len(records) > self.records_per_group:
|
| 57 |
+
records = random.sample(records, self.records_per_group)
|
| 58 |
+
|
| 59 |
+
# Observation
|
| 60 |
+
obs_data = records[0].observation_inline
|
| 61 |
+
if "state" in obs_data:
|
| 62 |
+
obs = list(obs_data["state"])
|
| 63 |
+
else:
|
| 64 |
+
obs = []
|
| 65 |
+
for v in obs_data.values():
|
| 66 |
+
if isinstance(v, list):
|
| 67 |
+
obs.extend(v)
|
| 68 |
+
elif isinstance(v, (int, float)):
|
| 69 |
+
obs.append(v)
|
| 70 |
+
obs = self._pad([float(x) for x in obs], self.max_obs_dim)
|
| 71 |
+
|
| 72 |
+
# Actions
|
| 73 |
+
actions = [self._pad(r.action_chunk.flat_values, self.max_act_dim) for r in records]
|
| 74 |
+
|
| 75 |
+
# Rewards
|
| 76 |
+
rewards = [r.reward.score for r in records]
|
| 77 |
+
|
| 78 |
+
# Language embedding
|
| 79 |
+
lang_emb = self.embeddings.get(group_id, np.zeros(768))
|
| 80 |
+
|
| 81 |
+
return {
|
| 82 |
+
"observation": torch.FloatTensor(obs),
|
| 83 |
+
"actions": torch.FloatTensor(actions),
|
| 84 |
+
"rewards": torch.FloatTensor(rewards),
|
| 85 |
+
"language": torch.FloatTensor(lang_emb) # NEW!
|
| 86 |
+
}
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def collate_fn(batch):
|
| 90 |
+
"""Collate with padding to max K in batch."""
|
| 91 |
+
max_k = max(b["actions"].shape[0] for b in batch)
|
| 92 |
+
batch_size = len(batch)
|
| 93 |
+
obs_dim = batch[0]["observation"].shape[0]
|
| 94 |
+
action_dim = batch[0]["actions"].shape[1]
|
| 95 |
+
lang_dim = batch[0]["language"].shape[0]
|
| 96 |
+
|
| 97 |
+
obs_batch = torch.stack([b["observation"] for b in batch])
|
| 98 |
+
actions_batch = torch.zeros(batch_size, max_k, action_dim)
|
| 99 |
+
rewards_batch = torch.zeros(batch_size, max_k)
|
| 100 |
+
lang_batch = torch.stack([b["language"] for b in batch]) # NEW!
|
| 101 |
+
|
| 102 |
+
for i, b in enumerate(batch):
|
| 103 |
+
k = b["actions"].shape[0]
|
| 104 |
+
actions_batch[i, :k] = b["actions"]
|
| 105 |
+
rewards_batch[i, :k] = b["rewards"]
|
| 106 |
+
|
| 107 |
+
return {
|
| 108 |
+
"observation": obs_batch,
|
| 109 |
+
"actions": actions_batch,
|
| 110 |
+
"rewards": rewards_batch,
|
| 111 |
+
"language": lang_batch # NEW!
|
| 112 |
+
}
|
| 113 |
+
|
| 114 |
+
|
| 115 |
+
def set_seed(seed: int):
|
| 116 |
+
random.seed(seed)
|
| 117 |
+
np.random.seed(seed)
|
| 118 |
+
torch.manual_seed(seed)
|
| 119 |
+
if torch.cuda.is_available():
|
| 120 |
+
torch.cuda.manual_seed_all(seed)
|
| 121 |
+
|
| 122 |
+
|
| 123 |
+
def get_cosine_schedule_with_warmup(optimizer, num_warmup_steps, num_training_steps):
|
| 124 |
+
"""Cosine schedule with linear warmup."""
|
| 125 |
+
def lr_lambda(current_step):
|
| 126 |
+
if current_step < num_warmup_steps:
|
| 127 |
+
return float(current_step) / float(max(1, num_warmup_steps))
|
| 128 |
+
progress = float(current_step - num_warmup_steps) / float(max(1, num_training_steps - num_warmup_steps))
|
| 129 |
+
return max(0.0, 0.5 * (1.0 + np.cos(np.pi * progress)))
|
| 130 |
+
|
| 131 |
+
return optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
|
| 132 |
+
|
| 133 |
+
|
| 134 |
+
def train_epoch(model: nn.Module, dataloader: DataLoader,
|
| 135 |
+
optimizer: optim.Optimizer, device: torch.device) -> float:
|
| 136 |
+
"""Train for one epoch."""
|
| 137 |
+
model.train()
|
| 138 |
+
total_loss = 0.0
|
| 139 |
+
num_batches = 0
|
| 140 |
+
|
| 141 |
+
for batch in dataloader:
|
| 142 |
+
obs = batch["observation"].to(device)
|
| 143 |
+
actions = batch["actions"].to(device)
|
| 144 |
+
rewards = batch["rewards"].to(device)
|
| 145 |
+
lang = batch["language"].to(device) # NEW!
|
| 146 |
+
|
| 147 |
+
optimizer.zero_grad()
|
| 148 |
+
|
| 149 |
+
scores = model(obs, actions, lang) # Pass language!
|
| 150 |
+
|
| 151 |
+
# Ranking loss
|
| 152 |
+
batch_size, K = rewards.shape
|
| 153 |
+
loss = 0.0
|
| 154 |
+
count = 0
|
| 155 |
+
|
| 156 |
+
for b in range(batch_size):
|
| 157 |
+
for i in range(K):
|
| 158 |
+
for j in range(K):
|
| 159 |
+
if i != j and rewards[b, i] != rewards[b, j]:
|
| 160 |
+
target = 1.0 if rewards[b, i] > rewards[b, j] else 0.0
|
| 161 |
+
pred = torch.sigmoid(scores[b, i, j])
|
| 162 |
+
loss += nn.functional.binary_cross_entropy(
|
| 163 |
+
pred.unsqueeze(0),
|
| 164 |
+
torch.tensor([target], device=device)
|
| 165 |
+
)
|
| 166 |
+
count += 1
|
| 167 |
+
|
| 168 |
+
if count > 0:
|
| 169 |
+
loss = loss / count
|
| 170 |
+
loss.backward()
|
| 171 |
+
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
|
| 172 |
+
optimizer.step()
|
| 173 |
+
total_loss += loss.item()
|
| 174 |
+
num_batches += 1
|
| 175 |
+
|
| 176 |
+
return total_loss / max(num_batches, 1)
|
| 177 |
+
|
| 178 |
+
|
| 179 |
+
def evaluate(model: nn.Module, dataloader: DataLoader, device: torch.device) -> dict:
|
| 180 |
+
"""Evaluate model."""
|
| 181 |
+
model.eval()
|
| 182 |
+
correct_top1 = 0
|
| 183 |
+
total_groups = 0
|
| 184 |
+
|
| 185 |
+
with torch.no_grad():
|
| 186 |
+
for batch in dataloader:
|
| 187 |
+
obs = batch["observation"].to(device)
|
| 188 |
+
actions = batch["actions"].to(device)
|
| 189 |
+
rewards = batch["rewards"].to(device)
|
| 190 |
+
lang = batch["language"].to(device) # NEW!
|
| 191 |
+
|
| 192 |
+
scores_matrix = model(obs, actions, lang) # Pass language!
|
| 193 |
+
|
| 194 |
+
batch_size, K = rewards.shape
|
| 195 |
+
for b in range(batch_size):
|
| 196 |
+
action_scores = scores_matrix[b].sum(dim=1).cpu().tolist()
|
| 197 |
+
selected = max(range(K), key=lambda i: action_scores[i])
|
| 198 |
+
best_utility = max(rewards[b].cpu().tolist())
|
| 199 |
+
|
| 200 |
+
if abs(rewards[b, selected].item() - best_utility) < 1e-6:
|
| 201 |
+
correct_top1 += 1
|
| 202 |
+
total_groups += 1
|
| 203 |
+
|
| 204 |
+
return {"top1_accuracy": correct_top1 / max(total_groups, 1)}
|
| 205 |
+
|
| 206 |
+
|
| 207 |
+
def main(argv: list[str] | None = None) -> int:
|
| 208 |
+
parser = argparse.ArgumentParser(description="Train DoVLA-Transformer with Language")
|
| 209 |
+
|
| 210 |
+
# Data
|
| 211 |
+
parser.add_argument("--dataset", type=Path, required=True)
|
| 212 |
+
parser.add_argument("--embeddings", type=Path, required=True) # NEW!
|
| 213 |
+
parser.add_argument("--out", type=Path, required=True)
|
| 214 |
+
|
| 215 |
+
# Architecture
|
| 216 |
+
parser.add_argument("--d-model", type=int, default=256)
|
| 217 |
+
parser.add_argument("--n-heads", type=int, default=8)
|
| 218 |
+
parser.add_argument("--n-layers", type=int, default=3)
|
| 219 |
+
parser.add_argument("--d-ff", type=int, default=1024)
|
| 220 |
+
|
| 221 |
+
# Training
|
| 222 |
+
parser.add_argument("--epochs", type=int, default=50)
|
| 223 |
+
parser.add_argument("--batch-size", type=int, default=16)
|
| 224 |
+
parser.add_argument("--lr", type=float, default=0.001)
|
| 225 |
+
parser.add_argument("--weight-decay", type=float, default=0.01)
|
| 226 |
+
parser.add_argument("--warmup-steps", type=int, default=500)
|
| 227 |
+
parser.add_argument("--seed", type=int, default=0)
|
| 228 |
+
parser.add_argument("--val-fraction", type=float, default=0.2)
|
| 229 |
+
parser.add_argument("--device", default="auto")
|
| 230 |
+
|
| 231 |
+
args = parser.parse_args(argv)
|
| 232 |
+
|
| 233 |
+
set_seed(args.seed)
|
| 234 |
+
args.out.mkdir(parents=True, exist_ok=True)
|
| 235 |
+
|
| 236 |
+
if args.device == "auto":
|
| 237 |
+
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
| 238 |
+
else:
|
| 239 |
+
device = torch.device(args.device)
|
| 240 |
+
|
| 241 |
+
print("=" * 70)
|
| 242 |
+
print("DoVLA-Transformer Training WITH LANGUAGE")
|
| 243 |
+
print("=" * 70)
|
| 244 |
+
print(f"Dataset: {args.dataset}")
|
| 245 |
+
print(f"Embeddings: {args.embeddings}")
|
| 246 |
+
print(f"Device: {device}")
|
| 247 |
+
print(f"Language dimension: 768")
|
| 248 |
+
print(f"Expected improvement: +8-11% (to 50-55%)")
|
| 249 |
+
print(f"Seed: {args.seed}")
|
| 250 |
+
print()
|
| 251 |
+
|
| 252 |
+
# Load embeddings
|
| 253 |
+
print("Loading instruction embeddings...")
|
| 254 |
+
with open(args.embeddings, 'rb') as f:
|
| 255 |
+
embeddings = pickle.load(f)
|
| 256 |
+
print(f"Loaded {len(embeddings)} embeddings")
|
| 257 |
+
print()
|
| 258 |
+
|
| 259 |
+
# Load data
|
| 260 |
+
print("Loading dataset...")
|
| 261 |
+
dataset = CILDataset(args.dataset)
|
| 262 |
+
all_groups = list(dataset.group_ids)
|
| 263 |
+
|
| 264 |
+
random.shuffle(all_groups)
|
| 265 |
+
split_idx = int(len(all_groups) * (1 - args.val_fraction))
|
| 266 |
+
train_groups = all_groups[:split_idx]
|
| 267 |
+
val_groups = all_groups[split_idx:]
|
| 268 |
+
|
| 269 |
+
print(f"Total: {len(all_groups)}, Train: {len(train_groups)}, Val: {len(val_groups)}")
|
| 270 |
+
print()
|
| 271 |
+
|
| 272 |
+
# Datasets
|
| 273 |
+
train_dataset = TransformerLangDataset(dataset, train_groups, embeddings)
|
| 274 |
+
val_dataset = TransformerLangDataset(dataset, val_groups, embeddings)
|
| 275 |
+
|
| 276 |
+
train_loader = DataLoader(train_dataset, batch_size=args.batch_size,
|
| 277 |
+
shuffle=True, num_workers=0, collate_fn=collate_fn)
|
| 278 |
+
val_loader = DataLoader(val_dataset, batch_size=args.batch_size,
|
| 279 |
+
shuffle=False, num_workers=0, collate_fn=collate_fn)
|
| 280 |
+
|
| 281 |
+
# Model
|
| 282 |
+
model = DoVLATransformer(
|
| 283 |
+
obs_dim=70,
|
| 284 |
+
action_dim=32,
|
| 285 |
+
lang_dim=768, # Enable language!
|
| 286 |
+
d_model=args.d_model,
|
| 287 |
+
n_heads=args.n_heads,
|
| 288 |
+
n_layers=args.n_layers,
|
| 289 |
+
d_ff=args.d_ff,
|
| 290 |
+
dropout=0.1
|
| 291 |
+
).to(device)
|
| 292 |
+
|
| 293 |
+
num_params = sum(p.numel() for p in model.parameters())
|
| 294 |
+
print(f"Model parameters: {num_params:,}")
|
| 295 |
+
print()
|
| 296 |
+
|
| 297 |
+
# Optimizer & Scheduler
|
| 298 |
+
optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
|
| 299 |
+
num_training_steps = len(train_loader) * args.epochs
|
| 300 |
+
scheduler = get_cosine_schedule_with_warmup(optimizer, args.warmup_steps, num_training_steps)
|
| 301 |
+
|
| 302 |
+
# Training
|
| 303 |
+
best_acc = 0.0
|
| 304 |
+
history = []
|
| 305 |
+
|
| 306 |
+
print("Starting training...")
|
| 307 |
+
print()
|
| 308 |
+
|
| 309 |
+
for epoch in range(args.epochs):
|
| 310 |
+
train_loss = train_epoch(model, train_loader, optimizer, device)
|
| 311 |
+
val_metrics = evaluate(model, val_loader, device)
|
| 312 |
+
scheduler.step()
|
| 313 |
+
|
| 314 |
+
val_acc = val_metrics["top1_accuracy"]
|
| 315 |
+
|
| 316 |
+
history.append({
|
| 317 |
+
"epoch": epoch + 1,
|
| 318 |
+
"train_loss": train_loss,
|
| 319 |
+
"val_top1_accuracy": val_acc,
|
| 320 |
+
"lr": scheduler.get_last_lr()[0]
|
| 321 |
+
})
|
| 322 |
+
|
| 323 |
+
print(f"Epoch {epoch+1:3d}/{args.epochs}: "
|
| 324 |
+
f"loss={train_loss:.4f}, val_top1={val_acc:.4f}, lr={scheduler.get_last_lr()[0]:.6f}")
|
| 325 |
+
|
| 326 |
+
if val_acc > best_acc:
|
| 327 |
+
best_acc = val_acc
|
| 328 |
+
torch.save({
|
| 329 |
+
"model_state_dict": model.state_dict(),
|
| 330 |
+
"epoch": epoch + 1,
|
| 331 |
+
"val_top1_accuracy": val_acc,
|
| 332 |
+
"args": vars(args)
|
| 333 |
+
}, args.out / "best.pt")
|
| 334 |
+
|
| 335 |
+
print()
|
| 336 |
+
print(f"✅ Training complete! Best val top-1: {best_acc:.4f}")
|
| 337 |
+
|
| 338 |
+
# Save
|
| 339 |
+
with open(args.out / "history.json", "w") as f:
|
| 340 |
+
json.dump(history, f, indent=2)
|
| 341 |
+
|
| 342 |
+
with open(args.out / "config.json", "w") as f:
|
| 343 |
+
json.dump({
|
| 344 |
+
"model": "DoVLA-Transformer-Language",
|
| 345 |
+
"language_dim": 768,
|
| 346 |
+
"architecture": {
|
| 347 |
+
"d_model": args.d_model,
|
| 348 |
+
"n_heads": args.n_heads,
|
| 349 |
+
"n_layers": args.n_layers,
|
| 350 |
+
"d_ff": args.d_ff
|
| 351 |
+
},
|
| 352 |
+
"training": {
|
| 353 |
+
"lr": args.lr,
|
| 354 |
+
"warmup_steps": args.warmup_steps,
|
| 355 |
+
"epochs": args.epochs,
|
| 356 |
+
"seed": args.seed
|
| 357 |
+
},
|
| 358 |
+
"results": {
|
| 359 |
+
"best_val_top1_accuracy": best_acc,
|
| 360 |
+
"num_parameters": num_params
|
| 361 |
+
}
|
| 362 |
+
}, f, indent=2)
|
| 363 |
+
|
| 364 |
+
return 0
|
| 365 |
+
|
| 366 |
+
|
| 367 |
+
if __name__ == "__main__":
|
| 368 |
+
sys.exit(main())
|