Complete standalone final deliverable and unaligned Q2 results
This commit is contained in:
+31
-5
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user