File size: 7,058 Bytes
59630ba | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 | """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()
|