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:
|
| 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
|
| 43 |
|
| 44 |
-
**Validation status.**
|
| 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;
|
| 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": "
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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"
|