Download shared/validate_minecraft_reference.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 7.06 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/shared/validate_minecraft_reference.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/shared/validate_minecraft_reference.py
-
curl -L -o validate_minecraft_reference.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/shared/validate_minecraft_reference.py
7.06 kB
| """Separate Oasis cache differences from encoder/decoder precision effects.""" | |
| import argparse | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| from protocol import read_frames, select_cases | |
| from score_case import load_oasis_vae | |
| SCALE = 0.07843137255 | |
| def differences(left, right): | |
| delta = left.float() - right.float() | |
| return { | |
| "max_absolute": delta.abs().flatten(1).amax(1).cpu().tolist(), | |
| "mse": delta.square().flatten(1).mean(1).cpu().tolist(), | |
| } | |
| def validate(manifest, data_root, feature_root, worldmem_root, vae_checkpoint, device): | |
| row = select_cases(manifest, limit=1)[0] | |
| relative = Path(row["relative_path"]) | |
| source_indices = np.asarray([700, 800, 1199]) | |
| feature_path = Path(feature_root) / relative.parent / (relative.stem + "_vae_feature.npy") | |
| metadata_path = feature_path.with_name(feature_path.stem + "_meta.json") | |
| metadata = json.loads(metadata_path.read_text()) | |
| array = np.load(feature_path, mmap_mode="r") | |
| expected_metadata = {"image_height": 360, "image_width": 640, "latent_channels": 16, | |
| "latent_height": 18, "latent_width": 32, "dtype": "float16", | |
| "scaling_factor": SCALE, "patch_size": 20} | |
| checks = {name: metadata.get(name) == value for name, value in expected_metadata.items()} | |
| checks.update({"array_dtype": array.dtype == np.float16, | |
| "array_shape": array.shape[1:] == (16, 18, 32), | |
| "metadata_frame_count": metadata["frame_count"] == len(array), | |
| "all_requested_source_frames": len(array) > int(source_indices[-1]), | |
| "manifest_target_range": int(row["target_start"]) <= 700 and int(row["target_end"]) > 1199, | |
| "metadata_source_suffix": Path(metadata["source_video"]).as_posix().endswith(relative.as_posix())}) | |
| if not all(checks.values()): | |
| raise ValueError(f"Canonical latent-cache metadata checks failed: {checks}") | |
| device = torch.device(device) | |
| raw = read_frames(Path(data_root) / relative, source_indices).to(device) | |
| cached = torch.from_numpy(np.array(array[source_indices], copy=True)).float().to(device) | |
| cached = cached.flatten(2).transpose(1, 2) | |
| vae = load_oasis_vae(worldmem_root, vae_checkpoint, device) | |
| autocast_enabled = device.type == "cuda" | |
| def encode(fp16): | |
| with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=fp16): | |
| # WorldMem.encode explicitly uses the posterior mean, never a sample. | |
| return vae.encode(raw * 2 - 1).mean * SCALE | |
| def decode(latent): | |
| with torch.autocast(device_type=device.type, dtype=torch.float16, enabled=autocast_enabled): | |
| return ((vae.decode(latent / SCALE) + 1) / 2).float().clamp(0, 1) | |
| native = encode(autocast_enabled) | |
| full_precision = encode(False) | |
| rounded = full_precision.to(torch.float16).float() | |
| native_rgb = decode(native) | |
| native_promoted_rgb = decode(native.float()) | |
| cached_rgb = decode(cached) | |
| fp32_latent_rgb = decode(full_precision) | |
| rounded_rgb = decode(rounded) | |
| rounding_error = (rounded - full_precision).abs() | |
| # Round-to-nearest fp16 has relative error <=2^-11 for normal values; | |
| # 2^-25 covers rounding of subnormal values. This bounds only the cast, | |
| # and provides no tolerance for neural-network decoding or cache lineage. | |
| rounding_bound = full_precision.abs() * (2.0 ** -11) + (2.0 ** -25) | |
| return { | |
| "device": str(device), "cuda_executed": device.type == "cuda", | |
| "source_clip": row["clip_id"], "source_indices": source_indices.tolist(), | |
| "vae_checkpoint": str(vae_checkpoint), "feature_path": str(feature_path), | |
| "cache_metadata": metadata, "cache_contract_checks": checks, | |
| "cache_lineage_limit": "Metadata records source/shape/scaling but does not identify encoder checkpoint, encoder arithmetic precision, or encoding batch size.", | |
| "raw_encoder_statistic": "posterior.mean (deterministic; no sampling)", | |
| "precision": {"native_scaled_latents": str(native.dtype), "cache_loader_latents": str(cached.dtype), | |
| "control_scaled_latents": str(full_precision.dtype), | |
| "decoder_autocast": "float16" if autocast_enabled else "disabled CPU diagnostic", | |
| "tf32_matmul": torch.backends.cuda.matmul.allow_tf32, | |
| "tf32_cudnn": torch.backends.cudnn.allow_tf32}, | |
| "latent_differences": { | |
| "cached_vs_native_encoder": differences(cached, native), | |
| "cached_vs_fp32_encoder": differences(cached, full_precision), | |
| "cached_vs_fp32_encoder_rounded_fp16": differences(cached, rounded), | |
| "native_encoder_vs_fp32_encoder": differences(native, full_precision), | |
| "pure_fp32_to_fp16_cast": differences(rounded, full_precision), | |
| }, | |
| "rgb_differences": { | |
| "cached_reference_vs_native_worldmem_reference": differences(cached_rgb, native_rgb), | |
| "cached_vs_native_with_same_fp32_scaled_latent_division": differences(cached_rgb, native_promoted_rgb), | |
| "native_decoder_fp16_vs_fp32_scaled_latent_division": differences(native_rgb, native_promoted_rgb), | |
| "pure_latent_quantization_same_decoder": differences(fp32_latent_rgb, rounded_rgb), | |
| }, | |
| "fp16_cast_bound": { | |
| "description": "abs(z)*2^-11 + 2^-25; cast only, not a VAE/cache acceptance threshold", | |
| "per_frame_max_bound": rounding_bound.flatten(1).amax(1).cpu().tolist(), | |
| "all_cast_errors_within_bound": bool((rounding_error <= rounding_bound).all()), | |
| }, | |
| "interpretation": "Numerical differences are reported without an equality assertion. Encoder precision, saved-latent quantization, decoder input arithmetic, and historical cache provenance are distinct effects.", | |
| "torch": torch.__version__, | |
| } | |
| def main(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--manifest", type=Path, required=True) | |
| parser.add_argument("--data-root", type=Path, required=True) | |
| parser.add_argument("--feature-root", type=Path) | |
| parser.add_argument("--worldmem-root", type=Path, required=True) | |
| parser.add_argument("--vae-checkpoint", type=Path, required=True) | |
| parser.add_argument("--output", type=Path, required=True) | |
| parser.add_argument("--device", default="cuda") | |
| args = parser.parse_args() | |
| # The fp32 encoder control must not silently use TensorFloat32 arithmetic. | |
| torch.backends.cuda.matmul.allow_tf32 = False | |
| torch.backends.cudnn.allow_tf32 = False | |
| result = validate(args.manifest, args.data_root, args.feature_root or args.data_root / "vae_features", | |
| args.worldmem_root, args.vae_checkpoint, args.device) | |
| args.output.parent.mkdir(parents=True, exist_ok=True) | |
| args.output.write_text(json.dumps(result, indent=2) + "\n") | |
| print(json.dumps(result, indent=2), flush=True) | |
| if __name__ == "__main__": | |
| main() | |