worldmem-baseline-evals / shared /validate_minecraft_reference.py
BonanDing's picture
Add isolated Minecraft and RE10K baseline evaluation suite
59630ba verified
Raw History Blame Contribute Delete
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(),
}
@torch.inference_mode()
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()