37 lines
1.4 KiB
Python
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()
|