anhtld commited on
Commit
20c251e
·
verified ·
1 Parent(s): adc02fa

Initial commit: DoVLA-CIL codebase (h=16 breakthrough) (part 2)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. scripts/slurm/generate_6task_h16.sbatch +87 -0
  2. scripts/slurm/generate_cil_array.sbatch +63 -0
  3. scripts/slurm/generate_embeddings.sbatch +28 -0
  4. scripts/slurm/horizon_sweep_pickcube.sbatch +88 -0
  5. scripts/slurm/install_smolvla_env.sbatch +104 -0
  6. scripts/slurm/make_maniskill_collection.sbatch +32 -0
  7. scripts/slurm/maniskill_lattice_debug.sbatch +111 -0
  8. scripts/slurm/maniskill_lattice_full.sbatch +80 -0
  9. scripts/slurm/maniskill_multitask_pilot.sbatch +63 -0
  10. scripts/slurm/phase_a1_generate_10k.sbatch +93 -0
  11. scripts/slurm/phase_a1_generate_10k_enhanced.sbatch +120 -0
  12. scripts/slurm/phase_a1_revised_enhanced.sbatch +65 -0
  13. scripts/slurm/phase_a1b_train_enhanced.sbatch +63 -0
  14. scripts/slurm/phase_a2_train_large_model.sbatch +59 -0
  15. scripts/slurm/phase_a3_eval_large_model.sbatch +50 -0
  16. scripts/slurm/phase_a4_hparam_sweep.sbatch +65 -0
  17. scripts/slurm/phase_a5_horizon_sweep.sbatch +63 -0
  18. scripts/slurm/phase_b_generate_12tasks.sbatch +109 -0
  19. scripts/slurm/phase_b_train_12tasks.sbatch +64 -0
  20. scripts/slurm/plan_c_generate_10k.sbatch +121 -0
  21. scripts/slurm/prepare_maniskill_baselines.sbatch +33 -0
  22. scripts/slurm/render_maniskill_multitask.sbatch +33 -0
  23. scripts/slurm/render_maniskill_observations.sbatch +49 -0
  24. scripts/slurm/run_external_vla_baseline.sbatch +51 -0
  25. scripts/slurm/run_scaling.sbatch +42 -0
  26. scripts/slurm/run_smolvla_cil_baseline.sbatch +48 -0
  27. scripts/slurm/smoke_smolvla_checkpoint.sbatch +42 -0
  28. scripts/slurm/train_attention_model.sbatch +62 -0
  29. scripts/slurm/train_dovla.sbatch +44 -0
  30. scripts/slurm/train_enhanced_model.sbatch +63 -0
  31. scripts/slurm/train_h16_policy.sbatch +54 -0
  32. scripts/slurm/train_hybrid_direct.sbatch +65 -0
  33. scripts/slurm/train_maniskill_baseline_array.sbatch +84 -0
  34. scripts/slurm/train_maniskill_collection_array.sbatch +114 -0
  35. scripts/slurm/train_maniskill_collection_cpu_array.sbatch +61 -0
  36. scripts/slurm/train_maniskill_debug.sbatch +86 -0
  37. scripts/slurm/train_maniskill_full_array.sbatch +74 -0
  38. scripts/slurm/train_maniskill_scaling_array.sbatch +82 -0
  39. scripts/slurm/train_maniskill_visual_array.sbatch +79 -0
  40. scripts/slurm/train_transformer.sbatch +72 -0
  41. scripts/slurm/train_transformer_lang.sbatch +68 -0
  42. scripts/smoke_full_pipeline.py +155 -0
  43. scripts/smoke_smolvla_checkpoint.py +169 -0
  44. scripts/smoke_test.sh +27 -0
  45. scripts/train_dovla.py +153 -0
  46. scripts/train_dovla_attention.py +324 -0
  47. scripts/train_dovla_enhanced.py +407 -0
  48. scripts/train_dovla_transformer.py +366 -0
  49. scripts/train_hybrid_direct.py +348 -0
  50. 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())