"""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()