"""Official Matrix demo and native-image controls; all scores are diagnostic.""" import argparse import json import os import subprocess import sys from pathlib import Path import numpy as np import torch from PIL import Image from diagnose_matrix_game2 import save_result from protocol import minecraft_actions, read_frames, select_cases from worldmem_adapter import map_actions ROOT = Path(__file__).resolve().parents[1] REPO = ROOT / 'Matrix-Game/Matrix-Game-2' MODES = ['official_demo', 'demo_actions', 'neutral', 'recorded', 'small_mouse'] def replay_controls(raw, mode, length): keyboard, mouse = map_actions(minecraft_actions(raw), 1.0) keys, turns = torch.zeros(length, 4), torch.zeros(length, 2) if mode != 'neutral': keys[:24], turns[:24] = keyboard, mouse * (0.1 if mode == 'small_mouse' else 1.0) return {'keyboard_condition': keys, 'mouse_condition': turns} 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('--dry-run', action='store_true') args = parser.parse_args() checkpoint = ROOT / 'checkpoints/matrix-game2' rows = select_cases(ROOT / 'deployment/manifests/minecraft_test_300.csv', limit=2) manifest = {'mode': args.mode, 'diagnostic_only': True, 'checkpoint': str(checkpoint), 'upstream_revision': '71c3cd7f741311f8100f6cf9cde942b6c1378d11', 'weight_policy': 'released raw distilled; original normalization; no optimizer or updates', 'expected_generator_parameters': 1619242304, 'expected_vae_parameters': 126892531, 'expected_clip_parameters': 1193014785, 'batch_size': 1, 'lora': None, 'lr': None, 'activation_checkpointing': False, 'steps': 3, 'shift': 5, 'cache_latents': 6, 'decoder': 'official compiled max-autotune-no-cudagraphs', 'official_demo': 'unchanged inference.py; supplied painting0000.png; seed0;150latents=597RGB', 'minecraft': {'anchor': 699, 'native_latents': 9, 'native_rgb': 33, 'scored_slice': [1, 25], 'targets': [700, 724], 'recorded_control_alignment': 'incoming700:724 plus9noops; small_mouse scales camera only', 'reference': 'rawRGB with official center crop and frame_process'}, 'score_caveat': 'Only recorded controls replay GT actions; other modes test stability, not benchmark accuracy', 'cases': ([{'image': 'demo_images/universal/0000.png', 'seed': 0}] if args.mode == 'official_demo' else [{'index': r['global_test_index'], 'path': r['relative_path'], 'seed': 42000000 + int(r['global_test_index'])} for r in rows]), 'output_frames_per_case': 597 if args.mode == 'official_demo' else 33} if args.dry_run: print(json.dumps(manifest, indent=2)) return args.output.mkdir(parents=True, exist_ok=True) (args.output / 'manifest.json').write_text(json.dumps(manifest, indent=2)) os.chdir(REPO) if args.mode == 'official_demo': subprocess.run([sys.executable, 'inference.py', '--checkpoint_path', str(checkpoint / 'base_distilled_model/base_distill.safetensors'), '--pretrained_model_path', str(checkpoint), '--img_path', 'demo_images/universal/0000.png', '--config_path', 'configs/inference_yaml/inference_universal.yaml', '--seed', '0', '--num_output_frames', '150', '--output_folder', str(args.output)], check=True) from torchvision.io import read_video video, _, _ = read_video(str(args.output / 'demo.mp4'), pts_unit='sec') assert len(video) == 597 Image.fromarray(torch.cat([video[t] for t in [0, 1, 2, 24, 100, 596]], dim=1).numpy()).save(args.output / 'samples.jpg') else: import inference as native native.set_seed(0) model_args = argparse.Namespace(config_path='configs/inference_yaml/inference_universal.yaml', checkpoint_path=str(checkpoint / 'base_distilled_model/base_distill.safetensors'), pretrained_model_path=str(checkpoint), num_output_frames=9, seed=0, img_path='', output_folder='') engine = native.InteractiveGameInference(model_args) manifest['generator_requires_grad_parameters'] = sum(p.numel() for p in engine.pipeline.generator.parameters() if p.requires_grad) manifest['inference_grad_policy'] = 'unchanged native no_grad; parameter flags unmodified' (args.output / 'manifest.json').write_text(json.dumps(manifest, indent=2)) original_inference, native_controls = engine.pipeline.inference, native.Bench_actions_universal captured = {} def capture(**kwargs): condition = kwargs['conditional_dict'] np.savez_compressed(case_dir / 'inputs.npz', initial_noise=kwargs['noise'].float().cpu().numpy(), keyboard=condition['keyboard_cond'][0].float().cpu().numpy(), mouse=condition['mouse_cond'][0].float().cpu().numpy()) videos = original_inference(**kwargs) captured['prediction'] = torch.cat(videos, dim=1)[0, 1:25].detach().float().cpu().add(1).div(2) return videos engine.pipeline.inference = capture for row in rows: case_dir = args.output / f"case{row['global_test_index']}" case_dir.mkdir(exist_ok=True) path = args.data_root / row['relative_path'] rgb = read_frames(path, torch.arange(699, 724)) images = [Image.fromarray(frame.mul(255).round().byte().permute(1, 2, 0).numpy()) for frame in rgb] images[0].save(case_dir / 'input.png') reference = torch.stack([engine.frame_process(engine._resizecrop(im, 352, 640)).add(1).div(2) for im in images[1:]]) with np.load(path.with_suffix('.npz')) as annotation: controls = replay_controls(annotation['actions'][700:724], args.mode, 33) native.Bench_actions_universal = native_controls if args.mode == 'demo_actions' else lambda length: controls model_args.img_path, model_args.output_folder = str(case_dir / 'input.png'), str(case_dir) engine.config['mode'] = 'universal' native.set_seed(42000000 + int(row['global_test_index'])) engine.generate_videos() save_result(case_dir, 'comparison', captured.pop('prediction'), reference) (args.output / 'complete.json').write_text(json.dumps({'mode': args.mode, 'complete': True})) print('NATIVE_CONTROL_COMPLETE', flush=True) if __name__ == '__main__': main()