Files

37 lines
1.4 KiB
Python

"""Load official unaligned data through Q2's real input entry, without training."""
from __future__ import annotations
import json
import sys
from pathlib import Path
import numpy as np
Q1 = Path(__file__).resolve().parent
ROOT = Q1.parents[1]
sys.path.insert(0, str(ROOT / "math" / "Q2"))
from data import load_official_splits # noqa: E402
def main() -> None:
source = ROOT / "E题数据" / "附件2-数据集特征文件" / "unaligned_50.pkl"
splits = load_official_splits(source, version="unaligned_50")
report = {}
for name, split in splits.items():
expected = {"text": (split.n, 50, 768), "audio": (split.n, 50, 74),
"vision": (split.n, 50, 35)}
actual = {m: x.shape for m, x in split.x.items()}
assert actual == expected
assert split.mask.shape == (split.n, 50, 3)
assert all(np.isfinite(x).all() for x in split.x.values())
assert split.class_y is not None and len(split.class_y) == split.n
report[name] = {"samples": split.n, "feature_shapes": {m: list(s) for m, s in actual.items()},
"mask_shape": list(split.mask.shape), "labels_preserved": True}
path = Q1 / "results" / "unified_adapter" / "q2_smoke.json"
path.write_text(json.dumps(report, indent=2), encoding="utf-8")
print(json.dumps(report, indent=2))
if __name__ == "__main__":
main()