Complete standalone final deliverable and unaligned Q2 results

This commit is contained in:
2026-09-25 22:22:37 +08:00
parent c6b018e5d0
commit adc9c2064b
267 changed files with 15479 additions and 7976 deletions
+31 -5
View File
@@ -2,6 +2,7 @@
from __future__ import annotations
import pickle
import sys
from dataclasses import dataclass
from pathlib import Path
from typing import Any
@@ -75,13 +76,18 @@ class SplitData:
regression_y: np.ndarray | None
ids: list[str]
groups: np.ndarray
alignment_audit: dict[str, Any] | None = None
@property
def n(self) -> int:
return len(self.ids)
def _extract_split(name: str, obj: dict[str, Any], with_labels: bool) -> SplitData:
def _extract_split(
name: str, obj: dict[str, Any], with_labels: bool,
mask_override: np.ndarray | None = None,
alignment_audit: dict[str, Any] | None = None,
) -> SplitData:
raw: dict[str, np.ndarray] = {}
masks = []
for modality in MODALITIES:
@@ -96,6 +102,11 @@ def _extract_split(name: str, obj: dict[str, Any], with_labels: bool) -> SplitDa
raw[modality] = arr
masks.append(observed)
mask = np.stack(masks, axis=-1)
if mask_override is not None:
override = np.asarray(mask_override, bool)
if override.shape != mask.shape:
raise ValueError(f"{name}: projected mask shape {override.shape} differs from {mask.shape}")
mask = override
ids = [_decode(v) for v in _one_dim(obj["id"])]
if len(ids) != len(mask):
raise ValueError(f"{name}: id count differs from feature count")
@@ -121,7 +132,7 @@ def _extract_split(name: str, obj: dict[str, Any], with_labels: bool) -> SplitDa
raise ValueError(f"{name}: polarity/regression label mismatch (sample, class, score): {examples}")
else:
class_y = regression_y = None
return SplitData(name, raw, mask, class_y, regression_y, ids, groups)
return SplitData(name, raw, mask, class_y, regression_y, ids, groups, alignment_audit)
def sample_group(sample_id: str) -> str:
@@ -129,12 +140,27 @@ def sample_group(sample_id: str) -> str:
return sample_id.split("$_$", 1)[0]
def load_official_splits(path: Path = ALIGNED_PATH) -> dict[str, SplitData]:
def load_official_splits(path: Path = ALIGNED_PATH, *, version: str = "aligned_50") -> dict[str, SplitData]:
if version not in {"aligned_50", "unaligned_50"}:
raise ValueError(f"unsupported feature version: {version}")
obj = restricted_load(path)
required = {"train", "valid", "test"}
if not isinstance(obj, dict) or not required.issubset(obj):
raise ValueError("aligned_50.pkl must contain train, valid, and test dictionaries")
splits = {name: _extract_split(name, obj[name], with_labels=True) for name in ("train", "valid", "test")}
raise ValueError(f"{path.name} must contain train, valid, and test dictionaries")
if version == "unaligned_50":
repo_dir = Path(__file__).resolve().parents[2]
if str(repo_dir) not in sys.path:
sys.path.insert(0, str(repo_dir))
from final.adapter import adapt_official_split
splits = {}
for name in ("train", "valid", "test"):
projected, mask, audit = adapt_official_split(obj[name])
fields = {**obj[name], **projected}
splits[name] = _extract_split(name, fields, with_labels=True,
mask_override=mask, alignment_audit=audit)
else:
splits = {name: _extract_split(name, obj[name], with_labels=True) for name in ("train", "valid", "test")}
del obj
return splits