Complete standalone final deliverable and unaligned Q2 results
This commit is contained in:
@@ -12,6 +12,7 @@ import hashlib
|
||||
import json
|
||||
import math
|
||||
import random
|
||||
import sys
|
||||
import time
|
||||
from collections import Counter, defaultdict
|
||||
from pathlib import Path
|
||||
@@ -305,7 +306,12 @@ def validation_aurc_bootstrap(
|
||||
return rows
|
||||
|
||||
|
||||
def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR) -> None:
|
||||
def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR,
|
||||
input_version: str = "aligned_50", batch_size: int = BATCH_SIZE) -> None:
|
||||
global BATCH_SIZE
|
||||
if batch_size < 1:
|
||||
raise ValueError("batch_size must be positive")
|
||||
BATCH_SIZE = batch_size
|
||||
if output_dir.exists() and any(output_dir.iterdir()):
|
||||
raise FileExistsError(f"refusing to overwrite non-empty result directory: {output_dir}")
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
@@ -313,12 +319,37 @@ def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR) -> None:
|
||||
if device.type == "cuda" and not torch.cuda.is_available():
|
||||
raise RuntimeError("CUDA was requested but is unavailable")
|
||||
|
||||
feature_path = ATTACHMENT2 / "aligned_50.pkl"
|
||||
raw_splits = load_splits(feature_path)
|
||||
if input_version not in {"aligned_50", "unaligned_50"}:
|
||||
raise ValueError(f"unsupported input version: {input_version}")
|
||||
feature_path = ATTACHMENT2 / f"{input_version}.pkl"
|
||||
if input_version == "unaligned_50":
|
||||
repo_root = Path(__file__).resolve().parents[3]
|
||||
if str(repo_root) not in sys.path:
|
||||
sys.path.insert(0, str(repo_root))
|
||||
from final.adapter import adapt_official_split
|
||||
from .data import _unpickle, _ids_and_targets
|
||||
|
||||
source = _unpickle(feature_path)
|
||||
raw_splits = {}
|
||||
adapter_audit = {}
|
||||
for name in ("train", "valid", "test"):
|
||||
arrays, mask, audit = adapt_official_split(source[name])
|
||||
ids, y_cls, y_reg = _ids_and_targets(source[name])
|
||||
raw_splits[name] = Split(tuple(arrays[m] for m in ("text", "audio", "vision")),
|
||||
mask, y_cls, y_reg, ids)
|
||||
adapter_audit[name] = audit
|
||||
del source
|
||||
groups = {name: {sid.split("$_$", 1)[0] for sid in split.ids}
|
||||
for name, split in raw_splits.items()}
|
||||
if any(groups[a] & groups[b] for a, b in (("train", "valid"), ("train", "test"), ("valid", "test"))):
|
||||
raise ValueError("official source-video groups overlap")
|
||||
else:
|
||||
raw_splits = load_splits(feature_path)
|
||||
adapter_audit = None
|
||||
train_raw, valid_raw, test_raw = raw_splits["train"], raw_splits["valid"], raw_splits["test"]
|
||||
stats = fit_robust_stats(train_raw)
|
||||
train, valid, test = (apply_robust_stats(s, stats) for s in (train_raw, valid_raw, test_raw))
|
||||
stats_path = output_dir / "aligned_robust_stats.npz"
|
||||
stats_path = output_dir / f"{input_version}_robust_stats.npz"
|
||||
stats.save(stats_path)
|
||||
dims = tuple(int(x.shape[-1]) for x in train.x)
|
||||
valid_scenarios = make_scenarios(valid, SCENARIO_SEED)
|
||||
@@ -432,7 +463,11 @@ def run(device_name: str = "auto", output_dir: Path = OUTPUT_DIR) -> None:
|
||||
"cuda_device": torch.cuda.get_device_name(0) if device.type == "cuda" else None,
|
||||
"feature_file": str(feature_path),
|
||||
"feature_sha256": sha256(feature_path),
|
||||
"representation": "official aligned_50 ordered positions; not Q1 physical-time bins",
|
||||
"representation": ("final.adapter relative-progress projection of official unaligned_50; not physical-time alignment"
|
||||
if input_version == "unaligned_50" else
|
||||
"official aligned_50 ordered positions; not Q1 physical-time bins"),
|
||||
"adapter": "final.adapter.adapt_official_split" if input_version == "unaligned_50" else None,
|
||||
"adapter_audit": adapter_audit,
|
||||
"train_valid_test_counts": {name: split.n for name, split in raw_splits.items()},
|
||||
"source_video_groups": {name: len({sid.split("$_$", 1)[0] for sid in split.ids}) for name, split in raw_splits.items()},
|
||||
"official_group_splits_disjoint": True,
|
||||
@@ -509,5 +544,8 @@ if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--device", default="auto", choices=("auto", "cuda", "cpu"))
|
||||
parser.add_argument("--output-dir", type=Path, default=OUTPUT_DIR)
|
||||
parser.add_argument("--input-version", choices=("aligned_50", "unaligned_50"), default="aligned_50")
|
||||
parser.add_argument("--batch-size", type=int, default=BATCH_SIZE)
|
||||
arguments = parser.parse_args()
|
||||
run(device_name=arguments.device, output_dir=arguments.output_dir)
|
||||
run(device_name=arguments.device, output_dir=arguments.output_dir,
|
||||
input_version=arguments.input_version, batch_size=arguments.batch_size)
|
||||
|
||||
Reference in New Issue
Block a user