"""Causal GF evaluation and diagnostics; teacher history is diagnostic only.""" import argparse import json import os from pathlib import Path import subprocess import sys import time import numpy as np import torch ROOT = Path(__file__).resolve().parents[1] sys.path[:0] = [str(ROOT / 'GeometryForcing'), str(ROOT / 'shared')] from protocol import case_seed, load_re10k_case, save_case, select_cases MODES = {'ar8_cfg4': (8, 4.0, False), 'ar12_cfg4': (12, 4.0, False), 'ar12_cfg1': (12, 1.0, False), 'teacher12_cfg4': (12, 4.0, True)} @torch.inference_mode() def generate(model, case, window, guidance, teacher, seed, frames=32): device = model.device torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) model.generator = torch.Generator(device=device).manual_seed(seed) history = model._normalize_x(case['rgb'][101-window:100].to(device)[None]) predictions, sampler_states, global_states = [], [], [] for offset in range(frames): target = 100 + offset indices = torch.arange(target-window+1, target+1, device=device) if teacher: history = model._normalize_x(case['rgb'][target-window+1:target].to(device)[None]) context = torch.cat([history, torch.zeros_like(history[:, :1])], dim=1) # Keep stabilization masks identical when replacing past RGB in the teacher control. mask = torch.where(indices < 100, 1, 2)[None] mask[:, -1] = 0 conditions = model._pad_to_max_tokens(case['cameras'][target-window+1:target+1].to(device)[None]) sampler_states.append(model.generator.get_state()) global_states.append(torch.cuda.get_rng_state(device) if device.type == 'cuda' else torch.random.get_rng_state()) with torch.autocast('cuda', dtype=torch.float16, enabled=device.type == 'cuda'): sample, _ = model._sample_sequence(batch_size=1, length=window, context=context, context_mask=mask, conditions=conditions, history_guidance=guidance) next_frame = sample[:, -1:] history = torch.cat([history[:, 1:], next_frame], dim=1) predictions.append(model._unnormalize_x(next_frame)[0, 0].float().cpu()) print(json.dumps({'generated': offset+1, 'target': target, 'normalized_range': [float(next_frame.min()), float(next_frame.max())]}), flush=True) prediction = torch.stack(predictions) if prediction.shape != (frames, 3, 256, 256) or not torch.isfinite(prediction).all(): raise ValueError('Incomplete or non-finite causal GF prediction') return prediction.clamp(0, 1), torch.stack(sampler_states), torch.stack(global_states) def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument('--mode', choices=MODES, required=True) parser.add_argument('--data-root', type=Path, required=True) parser.add_argument('--output', type=Path, required=True) parser.add_argument('--scope', choices=['monitor8', 'final100'], default='monitor8') parser.add_argument('--limit', type=int, default=2) parser.add_argument('--frames', type=int, default=32) parser.add_argument('--rank', type=int, default=int(os.environ.get('SLURM_PROCID', 0))) parser.add_argument('--world-size', type=int, default=int(os.environ.get('SLURM_NTASKS', 1))) parser.add_argument('--dry-run', action='store_true') args = parser.parse_args() window, scale, teacher = MODES[args.mode] manifest_path = ROOT / 'deployment/manifests/re10k_real_trajectory_h100_f399_scopes_v1.csv' rows = select_cases(manifest_path, args.scope, args.rank, args.world_size, args.limit) manifest = {'mode': args.mode, 'diagnostic_only': teacher or args.scope != 'final100' or args.frames != 300, 'teacher_history': teacher, 'scope': args.scope, 'rank': args.rank, 'world_size': args.world_size, 'workflow': 'one target at a time; no keyframes or interpolation', 'active_slots': window, 'past_rgb_frames': window-1, 'physical_slots': 16, 'noise_padding_slots': 16-window, 'camera_policy': 'past and current target only; current camera repeated in noise-padding slots', 'history_policy': 'GT past diagnostic' if teacher else 'self-generated feedback after observation100', 'feedback_values': 'normalized and unclipped; clamp only the exported RGB', 'mask_policy': 'observed1/generated2/target0; teacher keeps the same masks and stabilization', 'guidance_scale': scale, 'stabilization': 0.02, 'steps': 50, 'ddim_eta': 0, 'weights': 'released state_dict; unchanged; no EMA substitution', 'parameters': 458818051, 'trainable_parameters': 0, 'batch_size': 1, 'optimizer': None, 'lr': None, 'lora': None, 'activation_checkpointing': False, 'precision': 'fp16 autocast; fp32 weights/cameras', 'scored_timeline': [100, 100+args.frames], 'master_seed': 0, 'full_loop_evaluation': args.frames == 300, 'torch': torch.__version__, 'initial_noise_shape': [1, 16, 3, 256, 256], 'target_noise_slot': window-1, 'rng_note': 'same case seeds; AR8 and AR12 have different target noise slots', 'cases': [{'scene': r['scene_id'], 'seed': case_seed(r, 0)} for r in rows]} if args.dry_run: print(json.dumps(manifest, indent=2)) return args.output = args.output.resolve() args.data_root = args.data_root.resolve() args.output.mkdir(parents=True, exist_ok=True) manifest_file = args.output / ('manifest.json' if args.world_size == 1 else f'manifest_rank{args.rank}.json') manifest_file.write_text(json.dumps(manifest, indent=2)) os.chdir(ROOT / 'GeometryForcing') from evaluation.worldmem_adapter import GeometryForcingAdapter from algorithms.dfot.history_guidance import HistoryGuidance adapter = GeometryForcingAdapter(ROOT / 'checkpoints/geometry-forcing/geometry_forcing_state_dict.ckpt') guidance = HistoryGuidance.stabilized_vanilla(scale, 0.02, timesteps=adapter.model.timesteps, visualize=False) manifest['checkpoint'] = adapter.provenance['checkpoint'] manifest['loaded_tensors'] = adapter.provenance['loaded_backbone_tensors'] manifest['actual_parameters'] = adapter.provenance['parameters'] manifest_file.write_text(json.dumps(manifest, indent=2)) for row in rows: case = load_re10k_case(row, args.data_root) seed = case_seed(row, 0) start = time.monotonic() prediction, sampler_states, global_states = generate(adapter.model, case, window, guidance, teacher, seed, args.frames) directory = save_case(args.output, row, prediction, {**manifest, 'seed': seed, 'generated_frames': args.frames, 'history_frames': 100, 'source_indices': case['source_indices'].tolist(), 'generation_seconds': time.monotonic()-start}) np.savez_compressed(directory / 'rng_states.npz', sampler=sampler_states.numpy(), global_rng=global_states.numpy()) del prediction, case torch.cuda.empty_cache() subprocess.run([sys.executable, str(ROOT / 'shared/score_case.py'), '--case-dir', str(directory), '--dataset', 're10k', '--data-root', str(args.data_root), '--worldmem-root', str(ROOT / 'references/worldmem')], check=True) complete_file = args.output / ('complete.json' if args.world_size == 1 else f'worker_{args.rank:02d}_complete.json') complete_file.write_text(json.dumps({'mode': args.mode, 'rank': args.rank, 'cases': len(rows), 'frames': args.frames})) print('CAUSAL_CONTROL_COMPLETE', flush=True) if __name__ == '__main__': main()