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