BonanDing commited on
Commit
03a4c98
·
verified ·
1 Parent(s): c3781c5

Add strictly past-only Geometry Forcing context and guidance diagnostics

Browse files
README.md CHANGED
@@ -18,7 +18,7 @@ Four independent model workspaces and conda environments for the existing Minecr
18
  | DecMem public Wan reference | `DecMem/` | `/proj/cvl/users/x_fahkh2/envs/decmem-eval` | Minecraft: 600 observed + 500 generated frames |
19
  | Matrix-Game 2.0 | `Matrix-Game/Matrix-Game-2/` | `/proj/cvl/users/x_fahkh2/envs/matrix-game2-eval` | Minecraft: 600 observed + 500 generated frames |
20
  | LIVE | `LIVE/` | `/proj/cvl/users/x_fahkh2/envs/live-eval` | RE10K: 100 observed + 300 generated frames |
21
- | Geometry Forcing | `GeometryForcing/` | `/proj/cvl/users/x_fahkh2/envs/geometry-forcing-eval` | RE10K: 100 observed + 300 generated frames |
22
 
23
  The Berzelius checkout location is `/proj/cvl/users/x_fahkh2/worldmem-baseline-evals`. Submit jobs from Berzelius Hopper. First, `deployment/setup_environment_berzelius.sh` creates all four environments under `/proj/cvl/users/x_fahkh2/envs/`, builds required CUDA extensions, downloads checkpoints, validates data, and caches metric weights. It installs the official pinned Miniforge 24.7.1-2 manager under `/proj/cvl/users/x_fahkh2/envs/baseline-miniforge3` and restricts conda environment creation to the `conda-forge` channel. It uses that manager by its absolute path, without a Miniforge module or shell-profile changes. It uses the dedicated `berzelius-hopper-cpu` partition, 14 CPU cores, no GPUs, and an eight-hour limit. It runs no model inference. Setup logs are written under `slurm_logs/`, and records under `artifacts/environment_setup_JOBID/`. All setup, smoke and full evaluation jobs set both `TMPDIR` and `TRITON_CACHE_DIR` to `/proj/cvl/users/x_fahkh2/worldmem-baseline-evals/cache/`, as required by the project.
24
 
@@ -29,7 +29,9 @@ After setup completes successfully, each method has its own smoke job and full e
29
  | DecMem | `deployment/smoke_decmem_1h200.sh` | `deployment/evaluate_decmem_8h200.sh` |
30
  | Matrix-Game 2.0 | `deployment/smoke_matrix_game2_1h200.sh` | `deployment/evaluate_matrix_game2_8h200.sh` |
31
  | LIVE | `deployment/smoke_live_1h200.sh` | `deployment/evaluate_live_8h200.sh` |
32
- | Geometry Forcing | `deployment/smoke_geometry_forcing_1h200.sh` | `deployment/evaluate_geometry_forcing_8h200.sh` |
 
 
33
 
34
  Run smoke jobs separately and inspect prediction quality before the corresponding full evaluation. Each smoke uses one real-data case, runs generation and common scoring, then checks aggregation. LIVE generates 104 frames; Geometry Forcing generates the complete 300-frame loop to exercise both keyframe batches and interpolation. Minecraft tests generate 20 DecMem or 24 Matrix frames. Smoke logs are under the checkout's `slurm_logs/`; full logs are under each method workspace's `slurm_logs/`. Results are under each method workspace's `outputs/`.
35
 
@@ -39,8 +41,8 @@ Expected existing dataset locations are `/proj/cvl/users/x_fahkh2/WorldMem_Repro
39
 
40
  **Protocol details.** Minecraft observes source frames `[100,700)` and scores `[700,1200)` at 360×640, against the fixed Oasis-VAE reference cache. DecMem uses actions and poses with UniPC20/CFG5; Matrix uses calibrated action conditioning and its native three-step sampler. Temporal padding duplicates only observed frames, and the scored suffix contains exactly 500 frames. Matrix's incoming-to-outgoing action conversion is documented in its method manifest.
41
 
42
- RE10K follows `G0…G199,G199…G0`, including the repeated turnaround, with 100 observed and 300 generated frames at 256×256. LIVE uses its 18-step sampler with the requested 12-frame window. Geometry Forcing consumes only the last observed image (timeline 99), predicts 19 sparse keyframes over timeline 99–399, and interpolates the other frames with its native 16-slot backbone. Prediction guidance is 4 with stabilization 0.02; interpolation guidance is 1.5. Short camera batches repeat the final camera to fill unused model slots. This GF workflow is user-selected and is not a 12-slot comparison. Scene seeds follow `1_000_003 * metadata_index + 7`. The common RGB targets, AlexNet LPIPS, mean frame PSNR, SSIM and pooled FID are retained; these are not the paper's exact metric implementation. New GF outputs use the `geometry-forcing_native_re10k_` prefix.
43
 
44
- **Validation status.** The original adapters completed local GPU execution and scoring in job 4780, but subsequent visual checks found quality failures; those execution checks do not validate the revised workflows. The current GF adapter passed strict checkpoint loading and complete native-planner/camera-padding checks on CPU with a recording sampler. Actual GPU generation and visual quality require the new 300-frame smoke before the full run. The full GF job also repeats a one-case numerical preflight before its eight workers. Context budgets and method differences are recorded in `deployment/context_budget_audit.json`.
45
 
46
  Upstream revisions are recorded in `upstream_sources.json`. Sources retain their original licenses and notices, described in [THIRD_PARTY.md](THIRD_PARTY.md). In particular, the frozen WorldMem reference is covered by S-Lab License 1.0 and is not covered by the other projects' permissive licenses.
 
18
  | DecMem public Wan reference | `DecMem/` | `/proj/cvl/users/x_fahkh2/envs/decmem-eval` | Minecraft: 600 observed + 500 generated frames |
19
  | Matrix-Game 2.0 | `Matrix-Game/Matrix-Game-2/` | `/proj/cvl/users/x_fahkh2/envs/matrix-game2-eval` | Minecraft: 600 observed + 500 generated frames |
20
  | LIVE | `LIVE/` | `/proj/cvl/users/x_fahkh2/envs/live-eval` | RE10K: 100 observed + 300 generated frames |
21
+ | Geometry Forcing | `GeometryForcing/` | `/proj/cvl/users/x_fahkh2/envs/geometry-forcing-eval` | RE10K: causal diagnostics; native interpolation is an offline control |
22
 
23
  The Berzelius checkout location is `/proj/cvl/users/x_fahkh2/worldmem-baseline-evals`. Submit jobs from Berzelius Hopper. First, `deployment/setup_environment_berzelius.sh` creates all four environments under `/proj/cvl/users/x_fahkh2/envs/`, builds required CUDA extensions, downloads checkpoints, validates data, and caches metric weights. It installs the official pinned Miniforge 24.7.1-2 manager under `/proj/cvl/users/x_fahkh2/envs/baseline-miniforge3` and restricts conda environment creation to the `conda-forge` channel. It uses that manager by its absolute path, without a Miniforge module or shell-profile changes. It uses the dedicated `berzelius-hopper-cpu` partition, 14 CPU cores, no GPUs, and an eight-hour limit. It runs no model inference. Setup logs are written under `slurm_logs/`, and records under `artifacts/environment_setup_JOBID/`. All setup, smoke and full evaluation jobs set both `TMPDIR` and `TRITON_CACHE_DIR` to `/proj/cvl/users/x_fahkh2/worldmem-baseline-evals/cache/`, as required by the project.
24
 
 
29
  | DecMem | `deployment/smoke_decmem_1h200.sh` | `deployment/evaluate_decmem_8h200.sh` |
30
  | Matrix-Game 2.0 | `deployment/smoke_matrix_game2_1h200.sh` | `deployment/evaluate_matrix_game2_8h200.sh` |
31
  | LIVE | `deployment/smoke_live_1h200.sh` | `deployment/evaluate_live_8h200.sh` |
32
+ | Geometry Forcing (offline control) | `deployment/smoke_geometry_forcing_1h200.sh` | `deployment/evaluate_geometry_forcing_8h200.sh` |
33
+
34
+ For the strictly autoregressive comparison, use `deployment/geometry_causal_controls_1h200.sh`. Its four separate tasks test the previous 8-slot/guidance-4 setting, 12 slots/guidance 4, 12 slots/guidance 1, and a diagnostic that replaces generated history with ground-truth history while retaining the same noise masks. Each uses two monitor scenes and 32 future frames. The three autoregressive tasks use only past RGB and poses through the current target; unused model slots contain noise and copies of the current camera. The ground-truth-history task is not eligible as a benchmark result. These short tests must be reviewed before preparing a full causal evaluation.
35
 
36
  Run smoke jobs separately and inspect prediction quality before the corresponding full evaluation. Each smoke uses one real-data case, runs generation and common scoring, then checks aggregation. LIVE generates 104 frames; Geometry Forcing generates the complete 300-frame loop to exercise both keyframe batches and interpolation. Minecraft tests generate 20 DecMem or 24 Matrix frames. Smoke logs are under the checkout's `slurm_logs/`; full logs are under each method workspace's `slurm_logs/`. Results are under each method workspace's `outputs/`.
37
 
 
41
 
42
  **Protocol details.** Minecraft observes source frames `[100,700)` and scores `[700,1200)` at 360×640, against the fixed Oasis-VAE reference cache. DecMem uses actions and poses with UniPC20/CFG5; Matrix uses calibrated action conditioning and its native three-step sampler. Temporal padding duplicates only observed frames, and the scored suffix contains exactly 500 frames. Matrix's incoming-to-outgoing action conversion is documented in its method manifest.
43
 
44
+ RE10K follows `G0…G199,G199…G0`, including the repeated turnaround, with 100 observed and 300 generated frames at 256×256. LIVE uses its 18-step sampler with the requested 12-frame window. Geometry Forcing consumes only the last observed image (timeline 99), predicts 19 sparse keyframes over timeline 99–399, and interpolates the other frames with its native 16-slot backbone. Prediction guidance is 4 with stabilization 0.02; interpolation guidance is 1.5. Short camera batches repeat the final camera to fill unused model slots. This offline GF workflow is not a 12-slot comparison or a strictly autoregressive evaluation. Scene seeds follow `1_000_003 * metadata_index + 7`. The common RGB targets, AlexNet LPIPS, mean frame PSNR, SSIM and pooled FID are retained; these are not the paper's exact metric implementation. New GF outputs use the `geometry-forcing_native_re10k_` prefix.
45
 
46
+ **Validation status.** Native GF smoke 85127 completed 300 frames and scoring, but its interpolation uses later generated keyframes and is not a fair past-only autoregressive comparison. Retain that workflow as an offline diagnostic. The new causal controls have separate outputs; GPU image quality and stability remain to be established. Context budgets and method differences are recorded in `deployment/context_budget_audit.json`.
47
 
48
  Upstream revisions are recorded in `upstream_sources.json`. Sources retain their original licenses and notices, described in [THIRD_PARTY.md](THIRD_PARTY.md). In particular, the frozen WorldMem reference is covered by S-Lab License 1.0 and is not covered by the other projects' permissive licenses.
deployment/context_budget_audit.json CHANGED
@@ -436,7 +436,7 @@
436
  "gpu_image_quality_pending": true
437
  }
438
  },
439
- "scope": "LIVE12 retained; GFupdated to user-selectednativeworkflowoncommonloop; Minecraftmethodsunchanged",
440
  "geometry_forcing_previous_dense": {
441
  "audit_date": "2026-09-11",
442
  "method": "Geometry Forcing",
@@ -669,5 +669,39 @@
669
  ],
670
  "recommendation": "Retain current GF results explicitly labeled local8/no-memory. A GF active12 rerun is feasible for equal total frame budget, but is a separate experimental variant requiring authorized adapter changes and a smoke run; no change is made by this audit."
671
  },
672
- "geometry_forcing_budget_exception": "User selected native16-slotworkflow onourloop instead of12-slot comparison."
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
673
  }
 
436
  "gpu_image_quality_pending": true
437
  }
438
  },
439
+ "scope": "LIVE12 retained; GF nativeworkflow classifiedoffline; causal12controls prepared; Minecraftmethodsunchanged",
440
  "geometry_forcing_previous_dense": {
441
  "audit_date": "2026-09-11",
442
  "method": "Geometry Forcing",
 
669
  ],
670
  "recommendation": "Retain current GF results explicitly labeled local8/no-memory. A GF active12 rerun is feasible for equal total frame budget, but is a separate experimental variant requiring authorized adapter changes and a smoke run; no change is made by this audit."
671
  },
672
+ "geometry_forcing_budget_exception": "Previous native16-slot run retained only as an offline diagnostic; new past-only tests target12active slots.",
673
+ "geometry_forcing_native_classification": "Offline diagnostic only: interpolation uses later generated keyframes; not an eligible strict autoregressive baseline.",
674
+ "geometry_forcing_causal_controls": {
675
+ "scope": "GF early-collapse diagnostics, not full benchmark",
676
+ "modes": [
677
+ "ar8_cfg4",
678
+ "ar12_cfg4",
679
+ "ar12_cfg1",
680
+ "teacher12_cfg4"
681
+ ],
682
+ "dataset": "RE10K monitor8 first two fixed scenes",
683
+ "frames": 32,
684
+ "scored_timeline": [
685
+ 100,
686
+ 132
687
+ ],
688
+ "primary_active_budget": 12,
689
+ "physical_backbone_slots": 16,
690
+ "causal_conditioning": "RGB strictly beforetarget; cameras only through currenttarget; currentcamera repeated fornoise-padding",
691
+ "teacher_control": "Only pastRGB replacedwithGT; same masks andstabilization, diagnostic-only",
692
+ "interpolation": "nevercalled",
693
+ "frozen": true,
694
+ "parameters": 458818051,
695
+ "optimizer": null,
696
+ "lr": null,
697
+ "lora": null,
698
+ "sampling": "DDIM50eta0; stabilization.02; originalstate_dict",
699
+ "rng": "private/globalstatesrecorded;12slotCFGpair hasthe same targetnoise;8vs12changes targetslot",
700
+ "cpu_validation_files": [
701
+ "cli_validation.json",
702
+ "causality_validation.json",
703
+ "checkpoint_interface_validation.json"
704
+ ],
705
+ "gpu_validation": "pendingBerzeliussubmission"
706
+ }
707
  }
deployment/geometry_causal_controls.py ADDED
@@ -0,0 +1,109 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Past-only GF diagnostics; the teacher-history arm is not a benchmark result."""
2
+ import argparse
3
+ import json
4
+ import os
5
+ from pathlib import Path
6
+ import subprocess
7
+ import sys
8
+ import time
9
+ import numpy as np
10
+ import torch
11
+
12
+ ROOT = Path(__file__).resolve().parents[1]
13
+ sys.path[:0] = [str(ROOT / 'GeometryForcing'), str(ROOT / 'shared')]
14
+ from protocol import case_seed, load_re10k_case, save_case, select_cases
15
+
16
+ MODES = {'ar8_cfg4': (8, 4.0, False), 'ar12_cfg4': (12, 4.0, False),
17
+ 'ar12_cfg1': (12, 1.0, False), 'teacher12_cfg4': (12, 4.0, True)}
18
+
19
+ @torch.inference_mode()
20
+ def generate(model, case, window, guidance, teacher, seed, frames=32):
21
+ device = model.device
22
+ torch.manual_seed(seed)
23
+ torch.cuda.manual_seed_all(seed)
24
+ model.generator = torch.Generator(device=device).manual_seed(seed)
25
+ history = model._normalize_x(case['rgb'][101-window:100].to(device)[None])
26
+ predictions, sampler_states, global_states = [], [], []
27
+ for offset in range(frames):
28
+ target = 100 + offset
29
+ indices = torch.arange(target-window+1, target+1, device=device)
30
+ if teacher:
31
+ history = model._normalize_x(case['rgb'][target-window+1:target].to(device)[None])
32
+ context = torch.cat([history, torch.zeros_like(history[:, :1])], dim=1)
33
+ # Keep stabilization masks identical when replacing past RGB in the teacher control.
34
+ mask = torch.where(indices < 100, 1, 2)[None]
35
+ mask[:, -1] = 0
36
+ conditions = model._pad_to_max_tokens(case['cameras'][target-window+1:target+1].to(device)[None])
37
+ sampler_states.append(model.generator.get_state())
38
+ global_states.append(torch.cuda.get_rng_state(device) if device.type == 'cuda' else torch.random.get_rng_state())
39
+ with torch.autocast('cuda', dtype=torch.float16, enabled=device.type == 'cuda'):
40
+ sample, _ = model._sample_sequence(batch_size=1, length=window, context=context,
41
+ context_mask=mask, conditions=conditions, history_guidance=guidance)
42
+ next_frame = sample[:, -1:]
43
+ history = torch.cat([history[:, 1:], next_frame], dim=1)
44
+ predictions.append(model._unnormalize_x(next_frame)[0, 0].float().cpu())
45
+ print(json.dumps({'generated': offset+1, 'target': target,
46
+ 'normalized_range': [float(next_frame.min()), float(next_frame.max())]}), flush=True)
47
+ prediction = torch.stack(predictions)
48
+ if prediction.shape != (frames, 3, 256, 256) or not torch.isfinite(prediction).all():
49
+ raise ValueError('Incomplete or non-finite causal GF prediction')
50
+ return prediction.clamp(0, 1), torch.stack(sampler_states), torch.stack(global_states)
51
+
52
+ def main():
53
+ parser = argparse.ArgumentParser(description=__doc__)
54
+ parser.add_argument('--mode', choices=MODES, required=True)
55
+ parser.add_argument('--data-root', type=Path, required=True)
56
+ parser.add_argument('--output', type=Path, required=True)
57
+ parser.add_argument('--dry-run', action='store_true')
58
+ args = parser.parse_args()
59
+ window, scale, teacher = MODES[args.mode]
60
+ manifest_path = ROOT / 'deployment/manifests/re10k_real_trajectory_h100_f399_scopes_v1.csv'
61
+ rows = select_cases(manifest_path, 'monitor8', limit=2)
62
+ manifest = {'mode': args.mode, 'diagnostic_only': True, 'teacher_history': teacher,
63
+ 'workflow': 'one target at a time; no keyframes or interpolation', 'active_slots': window,
64
+ 'past_rgb_frames': window-1, 'physical_slots': 16, 'noise_padding_slots': 16-window,
65
+ 'camera_policy': 'past and current target only; current camera repeated in noise-padding slots',
66
+ 'history_policy': 'GT past diagnostic' if teacher else 'self-generated feedback after observation100',
67
+ 'feedback_values': 'normalized and unclipped; clamp only the exported RGB',
68
+ 'mask_policy': 'observed1/generated2/target0; teacher keeps the same masks and stabilization',
69
+ 'guidance_scale': scale, 'stabilization': 0.02, 'steps': 50, 'ddim_eta': 0,
70
+ 'weights': 'released state_dict; unchanged; no EMA substitution', 'parameters': 458818051,
71
+ 'trainable_parameters': 0, 'batch_size': 1, 'optimizer': None, 'lr': None, 'lora': None,
72
+ 'activation_checkpointing': False, 'precision': 'fp16 autocast; fp32 weights/cameras',
73
+ 'scored_timeline': [100, 132], 'master_seed': 0, 'full_loop_evaluation': False,
74
+ 'torch': torch.__version__, 'initial_noise_shape': [1, 16, 3, 256, 256], 'target_noise_slot': window-1,
75
+ 'rng_note': 'same case seeds; AR8 and AR12 have different target noise slots',
76
+ 'cases': [{'scene': r['scene_id'], 'seed': case_seed(r, 0)} for r in rows]}
77
+ if args.dry_run:
78
+ print(json.dumps(manifest, indent=2))
79
+ return
80
+ args.output = args.output.resolve()
81
+ args.data_root = args.data_root.resolve()
82
+ args.output.mkdir(parents=True, exist_ok=True)
83
+ (args.output / 'manifest.json').write_text(json.dumps(manifest, indent=2))
84
+ os.chdir(ROOT / 'GeometryForcing')
85
+ from evaluation.worldmem_adapter import GeometryForcingAdapter
86
+ from algorithms.dfot.history_guidance import HistoryGuidance
87
+ adapter = GeometryForcingAdapter(ROOT / 'checkpoints/geometry-forcing/geometry_forcing_state_dict.ckpt')
88
+ guidance = HistoryGuidance.stabilized_vanilla(scale, 0.02, timesteps=adapter.model.timesteps, visualize=False)
89
+ manifest['checkpoint'] = adapter.provenance['checkpoint']
90
+ manifest['loaded_tensors'] = adapter.provenance['loaded_backbone_tensors']
91
+ manifest['actual_parameters'] = adapter.provenance['parameters']
92
+ (args.output / 'manifest.json').write_text(json.dumps(manifest, indent=2))
93
+ for row in rows:
94
+ case = load_re10k_case(row, args.data_root)
95
+ seed = case_seed(row, 0)
96
+ start = time.monotonic()
97
+ prediction, sampler_states, global_states = generate(adapter.model, case, window, guidance, teacher, seed)
98
+ directory = save_case(args.output, row, prediction, {**manifest, 'seed': seed, 'generated_frames': 32,
99
+ 'history_frames': 100, 'source_indices': case['source_indices'].tolist(), 'generation_seconds': time.monotonic()-start})
100
+ np.savez_compressed(directory / 'rng_states.npz', sampler=sampler_states.numpy(), global_rng=global_states.numpy())
101
+ del prediction, case
102
+ torch.cuda.empty_cache()
103
+ subprocess.run([sys.executable, str(ROOT / 'shared/score_case.py'), '--case-dir', str(directory),
104
+ '--dataset', 're10k', '--data-root', str(args.data_root), '--worldmem-root', str(ROOT / 'references/worldmem')], check=True)
105
+ (args.output / 'complete.json').write_text(json.dumps({'mode': args.mode, 'cases': 2, 'frames': 32}))
106
+ print('CAUSAL_CONTROL_COMPLETE', flush=True)
107
+
108
+ if __name__ == '__main__':
109
+ main()
deployment/geometry_causal_controls_1h200.sh ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env bash
2
+ #SBATCH --job-name=geometry-forcing-causal-control
3
+ #SBATCH --array=0-3
4
+ #SBATCH --nodes=1
5
+ #SBATCH --ntasks=1
6
+ #SBATCH --cpus-per-task=14
7
+ #SBATCH --gres=gpu:1
8
+ #SBATCH --time=02:00:00
9
+ #SBATCH --account=berzelius-2026-196
10
+ #SBATCH --chdir=/proj/cvl/users/x_fahkh2/worldmem-baseline-evals
11
+ #SBATCH --output=slurm_logs/%x-%A_%a.out
12
+ #SBATCH --error=slurm_logs/%x-%A_%a.out
13
+ set -euo pipefail
14
+ module load buildenv-gcccuda/12.9.1-gcc11
15
+ ROOT=/proj/cvl/users/x_fahkh2/worldmem-baseline-evals
16
+ PYTHON=/proj/cvl/users/x_fahkh2/envs/geometry-forcing-eval/bin/python
17
+ DATA_ROOT=/proj/cvl/users/x_fahkh2/WorldMem_Repro/datasets/re10k_dfot/real-estate-10k
18
+ MODES=(ar8_cfg4 ar12_cfg4 ar12_cfg1 teacher12_cfg4)
19
+ MODE=${MODES[$SLURM_ARRAY_TASK_ID]}
20
+ OUTPUT=$ROOT/GeometryForcing/outputs/geometry_causal_${MODE}_job${SLURM_ARRAY_JOB_ID}_task${SLURM_ARRAY_TASK_ID}
21
+ export TMPDIR=$ROOT/cache
22
+ export TRITON_CACHE_DIR=$ROOT/cache
23
+ export HF_HOME=$ROOT/cache/huggingface
24
+ export TORCH_HOME=$ROOT/cache/torch
25
+ export XDG_CACHE_HOME=$ROOT/cache
26
+ export MPLCONFIGDIR=$ROOT/cache/matplotlib
27
+ export PYTHONPATH=$ROOT/GeometryForcing
28
+ export PYTHONNOUSERSITE=1 PYTHONUNBUFFERED=1 OMP_NUM_THREADS=4 MKL_NUM_THREADS=4
29
+ export WANDB_MODE=disabled HF_HUB_OFFLINE=1 TRANSFORMERS_OFFLINE=1
30
+ mkdir -p "$TMPDIR" "$OUTPUT"
31
+ "$PYTHON" "$ROOT/deployment/verify_node.py" --expected-gpus 1 --output "$OUTPUT"
32
+ cp "$0" "$OUTPUT/launcher.sh"
33
+ cp "$ROOT/deployment/geometry_causal_controls.py" "$OUTPUT/probe.py"
34
+ srun --ntasks=1 --gpus-per-task=1 --cpus-per-task=14 --cpu-bind=cores \
35
+ "$PYTHON" "$ROOT/deployment/geometry_causal_controls.py" \
36
+ --mode "$MODE" --data-root "$DATA_ROOT" --output "$OUTPUT"