648 lines
27 KiB
Python
648 lines
27 KiB
Python
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import math
|
|
import platform
|
|
import statistics
|
|
import time
|
|
from collections import defaultdict
|
|
from datetime import datetime, timezone
|
|
from pathlib import Path
|
|
from typing import Any, Mapping, Sequence
|
|
|
|
import numpy as np
|
|
import torch
|
|
from sklearn.model_selection import GroupKFold
|
|
|
|
from .compare_methods import (
|
|
_alignment_rows,
|
|
_batches,
|
|
_baseline_output,
|
|
_collect_representations,
|
|
_fit_learned_model,
|
|
_make_figures,
|
|
_save_alignment,
|
|
_seed_everything,
|
|
_write_csv,
|
|
)
|
|
from .experiment_data import (
|
|
FeatureSample,
|
|
collate_feature_samples,
|
|
fit_feature_stats,
|
|
load_feature_samples,
|
|
)
|
|
from .experiment_probes import (
|
|
run_shuffled_alignment_reconstruction_probe,
|
|
run_within_clip_temporal_retrieval_probe,
|
|
)
|
|
from .types import MODALITIES
|
|
|
|
|
|
VARIANTS = ("v2_a", "v2_b", "v2_c")
|
|
METRIC_COLUMNS = {
|
|
"alignment": (
|
|
"mvr",
|
|
"normalized_entropy",
|
|
"width80_source_positions",
|
|
"trajectory_span_fraction",
|
|
"c_row",
|
|
"c_far",
|
|
),
|
|
"retrieval": ("r_at_1", "r_at_3", "mase_slots", "exact_r_at_1"),
|
|
"reconstruction": (
|
|
"mae_aligned",
|
|
"mae_shuffled_mean",
|
|
"mae_shuffled_std",
|
|
"gain_align",
|
|
),
|
|
}
|
|
|
|
|
|
def _summarize(
|
|
rows: Sequence[Mapping[str, Any]], group_keys: Sequence[str], metrics: Sequence[str]
|
|
) -> list[dict[str, Any]]:
|
|
grouped: dict[tuple[Any, ...], list[Mapping[str, Any]]] = defaultdict(list)
|
|
for row in rows:
|
|
grouped[tuple(row[key] for key in group_keys)].append(row)
|
|
results = []
|
|
for key, values in grouped.items():
|
|
summary: dict[str, Any] = dict(zip(group_keys, key))
|
|
summary["n"] = len(values)
|
|
for metric in metrics:
|
|
numbers = [float(row[metric]) for row in values if row.get(metric) not in (None, "")]
|
|
numbers = [number for number in numbers if math.isfinite(number)]
|
|
if numbers:
|
|
summary[f"{metric}_mean"] = statistics.fmean(numbers)
|
|
summary[f"{metric}_std"] = statistics.stdev(numbers) if len(numbers) > 1 else 0.0
|
|
results.append(summary)
|
|
return results
|
|
|
|
|
|
def _comparison_table(summaries: Mapping[str, Sequence[Mapping[str, Any]]]) -> list[dict[str, Any]]:
|
|
alignment = {(row["method"], row["modality"]): row for row in summaries["alignment"]}
|
|
retrieval = {(row["method"], row["direction"]): row for row in summaries["retrieval"]}
|
|
reconstruction = {
|
|
(row["method"], row["target_modality"]): row for row in summaries["reconstruction"]
|
|
}
|
|
rows = []
|
|
for method in ("M1", "M2", "M3", "M4"):
|
|
row: dict[str, Any] = {"method": method}
|
|
for modality in MODALITIES:
|
|
metrics = alignment[(method, modality)]
|
|
for key in ("mvr", "normalized_entropy", "trajectory_span_fraction", "c_row", "c_far"):
|
|
row[f"{key}_{modality}"] = metrics.get(f"{key}_mean")
|
|
for direction in ("text_to_audio", "text_to_vision", "audio_to_vision"):
|
|
metrics = retrieval[(method, direction)]
|
|
for key in ("r_at_1", "r_at_3", "mase_slots", "exact_r_at_1"):
|
|
row[f"{key}_{direction}"] = metrics.get(f"{key}_mean")
|
|
for modality in MODALITIES:
|
|
metrics = reconstruction[(method, modality)]
|
|
for key in ("mae_aligned", "mae_shuffled_mean", "gain_align"):
|
|
row[f"{key}_{modality}"] = metrics.get(f"{key}_mean")
|
|
rows.append(row)
|
|
return rows
|
|
|
|
|
|
def _folds_from_baseline(
|
|
samples: Sequence[FeatureSample], baseline_splits: Path, requested_folds: int
|
|
) -> list[tuple[list[FeatureSample], list[FeatureSample], dict[str, Any]]]:
|
|
by_id = {sample.sample_id: sample for sample in samples}
|
|
if baseline_splits.is_file():
|
|
split_rows = json.loads(baseline_splits.read_text(encoding="utf-8"))
|
|
else:
|
|
groups = [sample.group_id for sample in samples]
|
|
splitter = GroupKFold(n_splits=requested_folds)
|
|
split_rows = []
|
|
for fold, (train, validation) in enumerate(
|
|
splitter.split(np.zeros(len(samples)), groups=groups), start=1
|
|
):
|
|
split_rows.append(
|
|
{
|
|
"fold": fold,
|
|
"train_sample_ids": [samples[index].sample_id for index in train],
|
|
"validation_sample_ids": [samples[index].sample_id for index in validation],
|
|
"train_video_ids": sorted({samples[index].group_id for index in train}),
|
|
"validation_video_ids": sorted(
|
|
{samples[index].group_id for index in validation}
|
|
),
|
|
}
|
|
)
|
|
if len(split_rows) != requested_folds:
|
|
raise ValueError(
|
|
f"baseline split file has {len(split_rows)} folds; expected {requested_folds}"
|
|
)
|
|
folds = []
|
|
seen_validation: list[str] = []
|
|
for row in split_rows:
|
|
train_ids = row["train_sample_ids"]
|
|
validation_ids = row["validation_sample_ids"]
|
|
if set(train_ids) & set(validation_ids):
|
|
raise ValueError(f"sample leakage in fold {row['fold']}")
|
|
train_samples = [by_id[sample_id] for sample_id in train_ids]
|
|
val_samples = [by_id[sample_id] for sample_id in validation_ids]
|
|
train_groups = {sample.group_id for sample in train_samples}
|
|
val_groups = {sample.group_id for sample in val_samples}
|
|
if train_groups & val_groups:
|
|
raise ValueError(f"video_id leakage in fold {row['fold']}")
|
|
seen_validation.extend(validation_ids)
|
|
folds.append((train_samples, val_samples, row))
|
|
if len(seen_validation) != len(samples) or set(seen_validation) != set(by_id):
|
|
raise ValueError("baseline folds do not cover the current complete sample manifest")
|
|
return folds
|
|
|
|
|
|
def _evaluate_fixed_methods(
|
|
args: argparse.Namespace,
|
|
samples: Sequence[FeatureSample],
|
|
folds: Sequence[tuple[list[FeatureSample], list[FeatureSample], dict[str, Any]]],
|
|
device: torch.device,
|
|
output_dir: Path,
|
|
) -> tuple[dict[str, list[dict[str, Any]]], dict[str, dict[str, Mapping[str, np.ndarray]]]]:
|
|
output_dir.mkdir(parents=True, exist_ok=True)
|
|
rows: dict[str, list[dict[str, Any]]] = {
|
|
"alignment": [],
|
|
"retrieval": [],
|
|
"reconstruction": [],
|
|
}
|
|
example_weights: dict[str, dict[str, Mapping[str, np.ndarray]]] = defaultdict(dict)
|
|
sample_by_id = {sample.sample_id: sample for sample in samples}
|
|
example_id = args.example_id if args.example_id in sample_by_id else samples[0].sample_id
|
|
|
|
for fold_index, (train_samples, val_samples, split_row) in enumerate(folds, start=1):
|
|
stats = fit_feature_stats(train_samples)
|
|
print(
|
|
f"[fixed fold {fold_index}/{len(folds)}] train={len(train_samples)} "
|
|
f"validation={len(val_samples)}; evaluating unchanged M1/M2",
|
|
flush=True,
|
|
)
|
|
for method in ("M1", "M2"):
|
|
train_aligned, _ = _collect_representations(
|
|
method,
|
|
train_samples,
|
|
stats,
|
|
device=device,
|
|
grid_size=args.grid_size,
|
|
batch_size=args.batch_size,
|
|
)
|
|
val_aligned, val_weights = _collect_representations(
|
|
method,
|
|
val_samples,
|
|
stats,
|
|
device=device,
|
|
grid_size=args.grid_size,
|
|
batch_size=args.batch_size,
|
|
)
|
|
combined = {**train_aligned, **val_aligned}
|
|
for batch_samples in _batches(
|
|
val_samples,
|
|
args.batch_size,
|
|
shuffle=False,
|
|
rng=np.random.default_rng(args.seeds[0] + fold_index),
|
|
):
|
|
sequences, durations, intervals = collate_feature_samples(
|
|
batch_samples, stats, device
|
|
)
|
|
output = _baseline_output(method, sequences, durations, intervals, args.grid_size)
|
|
rows["alignment"].extend(
|
|
_alignment_rows(
|
|
method,
|
|
"fixed",
|
|
fold_index,
|
|
batch_samples,
|
|
output,
|
|
device,
|
|
args.mvr_epsilon,
|
|
)
|
|
)
|
|
for sample in val_samples:
|
|
_save_alignment(output_dir, method, "fixed", sample, val_weights[sample.sample_id])
|
|
if sample.sample_id == example_id:
|
|
example_weights[sample.sample_id][method] = val_weights[sample.sample_id]
|
|
|
|
probe_seed = args.seeds[0] + fold_index * 100
|
|
rows["retrieval"].extend(
|
|
{"method": method, "seed": "fixed", "fold": fold_index, **row}
|
|
for row in run_within_clip_temporal_retrieval_probe(
|
|
[sample.sample_id for sample in train_samples],
|
|
[sample.sample_id for sample in val_samples],
|
|
combined,
|
|
device=device,
|
|
seed=probe_seed,
|
|
epochs=args.retrieval_probe_epochs,
|
|
batch_size=args.probe_batch_size,
|
|
tolerance=args.retrieval_tolerance,
|
|
top_k=3,
|
|
)
|
|
)
|
|
rows["reconstruction"].extend(
|
|
{"method": method, "seed": "fixed", "fold": fold_index, **row}
|
|
for row in run_shuffled_alignment_reconstruction_probe(
|
|
[sample.sample_id for sample in train_samples],
|
|
[sample.sample_id for sample in val_samples],
|
|
combined,
|
|
device=device,
|
|
seed=probe_seed + 1,
|
|
ratio=args.mask_ratio,
|
|
epochs=args.reconstruction_probe_epochs,
|
|
batch_size=args.probe_batch_size,
|
|
shuffle_repeats=args.shuffle_repeats,
|
|
)
|
|
)
|
|
|
|
for key, values in rows.items():
|
|
_write_csv(output_dir / f"{key}_metrics.csv", values)
|
|
(output_dir / "README.md").write_text(
|
|
"# Fixed M1/M2 reference probes\n\n"
|
|
"M1 and M2 are recomputed from the same saved video-grouped folds and unchanged. "
|
|
"The within-clip retrieval projection is fitted on each training fold; held-out candidates "
|
|
"come only from the same clip. The reconstruction decoder is trained on aligned training "
|
|
"representations, then compared with a control that shuffles the two non-target streams.\n",
|
|
encoding="utf-8",
|
|
)
|
|
return rows, example_weights
|
|
|
|
|
|
def _write_stage_summary(
|
|
stage_dir: Path,
|
|
fixed_rows: Mapping[str, Sequence[dict[str, Any]]],
|
|
learned_rows: Mapping[str, Sequence[dict[str, Any]]],
|
|
) -> dict[str, list[dict[str, Any]]]:
|
|
all_rows = {
|
|
key: [*fixed_rows[key], *learned_rows[key]]
|
|
for key in ("alignment", "retrieval", "reconstruction")
|
|
}
|
|
group_columns = {
|
|
"alignment": ("method", "modality"),
|
|
"retrieval": ("method", "direction"),
|
|
"reconstruction": ("method", "target_modality"),
|
|
}
|
|
summaries = {
|
|
key: _summarize(rows, group_columns[key], METRIC_COLUMNS[key])
|
|
for key, rows in all_rows.items()
|
|
}
|
|
for key, rows in all_rows.items():
|
|
_write_csv(stage_dir / f"{key}_metrics_with_fixed.csv", rows)
|
|
_write_csv(stage_dir / f"{key}_summary.csv", summaries[key])
|
|
_write_csv(stage_dir / "comparison_summary.csv", _comparison_table(summaries))
|
|
(stage_dir / "summary.json").write_text(
|
|
json.dumps(summaries, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
|
)
|
|
return summaries
|
|
|
|
|
|
def _run_variant(
|
|
variant: str,
|
|
args: argparse.Namespace,
|
|
samples: Sequence[FeatureSample],
|
|
folds: Sequence[tuple[list[FeatureSample], list[FeatureSample], dict[str, Any]]],
|
|
fixed_rows: Mapping[str, Sequence[dict[str, Any]]],
|
|
fixed_examples: Mapping[str, Mapping[str, Mapping[str, np.ndarray]]],
|
|
device: torch.device,
|
|
) -> dict[str, Any]:
|
|
start_time = time.time()
|
|
stage_dir = args.output_dir / variant
|
|
stage_dir.mkdir(parents=True, exist_ok=True)
|
|
(stage_dir / "splits.json").write_text(
|
|
json.dumps([row for _, _, row in folds], ensure_ascii=False, indent=2), encoding="utf-8"
|
|
)
|
|
learned_rows: dict[str, list[dict[str, Any]]] = {
|
|
"alignment": [],
|
|
"retrieval": [],
|
|
"reconstruction": [],
|
|
}
|
|
training_summary: list[dict[str, Any]] = []
|
|
training_history: list[dict[str, Any]] = []
|
|
example_weights: dict[str, dict[str, Mapping[str, np.ndarray]]] = defaultdict(dict)
|
|
sample_by_id = {sample.sample_id: sample for sample in samples}
|
|
example_id = args.example_id if args.example_id in sample_by_id else samples[0].sample_id
|
|
example_sample = sample_by_id[example_id]
|
|
if example_id in fixed_examples:
|
|
example_weights[example_id].update(fixed_examples[example_id])
|
|
|
|
for fold_index, (train_samples, val_samples, _) in enumerate(folds, start=1):
|
|
stats = fit_feature_stats(train_samples)
|
|
fold_dir = stage_dir / f"fold_{fold_index:02d}"
|
|
print(
|
|
f"[{variant} fold {fold_index}/{len(folds)}] training M3/M4 with "
|
|
f"{len(train_samples)} train and {len(val_samples)} validation clips",
|
|
flush=True,
|
|
)
|
|
for seed in args.seeds:
|
|
for method in ("M3", "M4"):
|
|
checkpoint = fold_dir / f"seed_{seed}" / f"{method}.pt"
|
|
model, info = _fit_learned_model(
|
|
method,
|
|
train_samples,
|
|
val_samples,
|
|
stats,
|
|
device=device,
|
|
seed=seed + fold_index * 1009,
|
|
grid_size=args.grid_size,
|
|
hidden_size=args.hidden_size,
|
|
heads=args.heads,
|
|
dropout=args.dropout,
|
|
batch_size=args.batch_size,
|
|
max_epochs=args.epochs,
|
|
patience=args.patience,
|
|
learning_rate=args.learning_rate,
|
|
checkpoint_path=checkpoint,
|
|
loss_variant=variant,
|
|
)
|
|
training_summary.append(
|
|
{
|
|
"loss_variant": variant,
|
|
"method": method,
|
|
"seed": seed,
|
|
"fold": fold_index,
|
|
"best_epoch": info["best_epoch"],
|
|
"best_validation_objective": info["best_validation_objective"],
|
|
**{
|
|
key: value
|
|
for key, value in info["best_validation_metrics"].items()
|
|
if key not in {"epoch", "train_total"}
|
|
},
|
|
"checkpoint": info["checkpoint"],
|
|
}
|
|
)
|
|
training_history.extend(
|
|
{
|
|
"loss_variant": variant,
|
|
"method": method,
|
|
"seed": seed,
|
|
"fold": fold_index,
|
|
**epoch,
|
|
}
|
|
for epoch in info["history"]
|
|
)
|
|
print(
|
|
f"[{variant} fold {fold_index}] {method} seed={seed} "
|
|
f"best_epoch={info['best_epoch']} "
|
|
f"val={info['best_validation_objective']:.4f} "
|
|
f"C_row(audio/vision)="
|
|
f"{info['best_validation_metrics']['validation_c_row_audio']:.3f}/"
|
|
f"{info['best_validation_metrics']['validation_c_row_vision']:.3f}",
|
|
flush=True,
|
|
)
|
|
train_aligned, _ = _collect_representations(
|
|
method,
|
|
train_samples,
|
|
stats,
|
|
device=device,
|
|
grid_size=args.grid_size,
|
|
batch_size=args.batch_size,
|
|
model=model,
|
|
)
|
|
val_aligned, val_weights = _collect_representations(
|
|
method,
|
|
val_samples,
|
|
stats,
|
|
device=device,
|
|
grid_size=args.grid_size,
|
|
batch_size=args.batch_size,
|
|
model=model,
|
|
)
|
|
combined = {**train_aligned, **val_aligned}
|
|
for batch_samples in _batches(
|
|
val_samples,
|
|
args.batch_size,
|
|
shuffle=False,
|
|
rng=np.random.default_rng(seed + fold_index),
|
|
):
|
|
sequences, durations, _ = collate_feature_samples(
|
|
batch_samples, stats, device
|
|
)
|
|
output = model(sequences)
|
|
learned_rows["alignment"].extend(
|
|
_alignment_rows(
|
|
method,
|
|
str(seed),
|
|
fold_index,
|
|
batch_samples,
|
|
output,
|
|
device,
|
|
args.mvr_epsilon,
|
|
)
|
|
)
|
|
for sample in val_samples:
|
|
_save_alignment(stage_dir, method, str(seed), sample, val_weights[sample.sample_id])
|
|
if sample.sample_id == example_id and seed == args.seeds[0]:
|
|
example_weights[sample.sample_id][method] = val_weights[sample.sample_id]
|
|
|
|
probe_seed = seed + fold_index * 100 + (3 if method == "M3" else 7)
|
|
learned_rows["retrieval"].extend(
|
|
{
|
|
"method": method,
|
|
"seed": seed,
|
|
"fold": fold_index,
|
|
**row,
|
|
}
|
|
for row in run_within_clip_temporal_retrieval_probe(
|
|
[sample.sample_id for sample in train_samples],
|
|
[sample.sample_id for sample in val_samples],
|
|
combined,
|
|
device=device,
|
|
seed=probe_seed,
|
|
epochs=args.retrieval_probe_epochs,
|
|
batch_size=args.probe_batch_size,
|
|
tolerance=args.retrieval_tolerance,
|
|
top_k=3,
|
|
)
|
|
)
|
|
learned_rows["reconstruction"].extend(
|
|
{
|
|
"method": method,
|
|
"seed": seed,
|
|
"fold": fold_index,
|
|
**row,
|
|
}
|
|
for row in run_shuffled_alignment_reconstruction_probe(
|
|
[sample.sample_id for sample in train_samples],
|
|
[sample.sample_id for sample in val_samples],
|
|
combined,
|
|
device=device,
|
|
seed=probe_seed + 1,
|
|
ratio=args.mask_ratio,
|
|
epochs=args.reconstruction_probe_epochs,
|
|
batch_size=args.probe_batch_size,
|
|
shuffle_repeats=args.shuffle_repeats,
|
|
)
|
|
)
|
|
del model
|
|
if device.type == "cuda":
|
|
torch.cuda.empty_cache()
|
|
|
|
_write_csv(stage_dir / "training_summary.csv", training_summary)
|
|
_write_csv(stage_dir / "training_history.csv", training_history)
|
|
for key, rows in learned_rows.items():
|
|
_write_csv(stage_dir / f"{key}_metrics_learned_only.csv", rows)
|
|
|
|
summaries = _write_stage_summary(stage_dir, fixed_rows, learned_rows)
|
|
_write_csv(stage_dir / "training_summary.csv", training_summary)
|
|
_write_csv(stage_dir / "training_history.csv", training_history)
|
|
if example_id in example_weights and set(example_weights[example_id]) == {"M1", "M2", "M3", "M4"}:
|
|
_make_figures(stage_dir, example_id, example_sample, example_weights[example_id], args.grid_size)
|
|
manifest = {
|
|
"variant": variant,
|
|
"loss_coefficients": {
|
|
"lambda_reconstruction": 1.0,
|
|
"lambda_contrastive": 1.0,
|
|
"lambda_monotonicity": 0.1,
|
|
"lambda_span": 5.0,
|
|
"lambda_diversity": 0.5 if variant in {"v2_b", "v2_c"} else 0.0,
|
|
"lambda_band": 10.0 if variant == "v2_c" else 0.0,
|
|
"coverage_floor": 0.7,
|
|
"diversity_slot_separation": 6,
|
|
"band_margin": 0.1,
|
|
},
|
|
"sample_count": len(samples),
|
|
"video_group_folds": len(folds),
|
|
"seeds": args.seeds,
|
|
"device": str(device),
|
|
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
|
|
"python": platform.python_version(),
|
|
"torch": torch.__version__,
|
|
"elapsed_seconds": time.time() - start_time,
|
|
"probes": {
|
|
"within_clip_retrieval_tolerance_slots": args.retrieval_tolerance,
|
|
"retrieval_top_k": 3,
|
|
"shuffled_reconstruction_repeats": args.shuffle_repeats,
|
|
"masked_block_ratio": args.mask_ratio,
|
|
},
|
|
"limits": [
|
|
"Retrieval projections are fitted on training-fold grid-slot positives; test candidates are restricted to the same held-out clip.",
|
|
"The reconstruction control shuffles the two non-target modality slot streams and preserves the target stream.",
|
|
"No human event timestamps are available, so temporal probes do not replace manual annotation.",
|
|
"Slot-regularization losses impose weak temporal structure and must be interpreted alongside the unregularized M1/M2 reference.",
|
|
],
|
|
}
|
|
(stage_dir / "run_manifest.json").write_text(
|
|
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
|
)
|
|
(stage_dir / "README.md").write_text(
|
|
f"# {variant} M3/M4 alignment variant\n\n"
|
|
"M1/M2 rows in the comparison files are fixed references recomputed on the original grouped folds. "
|
|
"Only M3/M4 training losses changed. See `training_history.csv` for per-epoch loss components and row-collapse scores. "
|
|
"`retrieval_metrics_learned_only.csv` restricts candidates to the same clip. "
|
|
"`reconstruction_metrics_learned_only.csv` contrasts aligned and shuffled non-target streams. "
|
|
"The experiment does not include RoPE, Gaussian bias, or latent-length changes.\n",
|
|
encoding="utf-8",
|
|
)
|
|
print(
|
|
f"[{variant} done] elapsed={manifest['elapsed_seconds']:.1f}s "
|
|
f"output={stage_dir}; summary rows={sum(len(rows) for rows in summaries.values())}",
|
|
flush=True,
|
|
)
|
|
return manifest
|
|
|
|
|
|
def run(args: argparse.Namespace) -> dict[str, Any]:
|
|
start_time = time.time()
|
|
_seed_everything(args.seeds[0])
|
|
if args.device == "auto":
|
|
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
|
else:
|
|
device = torch.device(args.device)
|
|
if device.type == "cuda" and not torch.cuda.is_available():
|
|
raise RuntimeError("CUDA was requested but is not available")
|
|
samples = load_feature_samples(args.feature_dir, args.manifest)
|
|
if len(samples) != args.expected_samples:
|
|
raise ValueError(f"expected {args.expected_samples} samples, found {len(samples)}")
|
|
folds = _folds_from_baseline(
|
|
samples, args.baseline_dir / "splits.json", args.folds
|
|
)
|
|
args.output_dir.mkdir(parents=True, exist_ok=True)
|
|
print(
|
|
f"[start] samples={len(samples)} folds={len(folds)} seeds={args.seeds} "
|
|
f"device={device} variants={','.join(VARIANTS)}",
|
|
flush=True,
|
|
)
|
|
fixed_rows, fixed_examples = _evaluate_fixed_methods(
|
|
args,
|
|
samples,
|
|
folds,
|
|
device,
|
|
args.output_dir / "fixed_baselines",
|
|
)
|
|
stage_manifests = []
|
|
for variant in VARIANTS:
|
|
stage_manifests.append(
|
|
_run_variant(
|
|
variant,
|
|
args,
|
|
samples,
|
|
folds,
|
|
fixed_rows,
|
|
fixed_examples,
|
|
device,
|
|
)
|
|
)
|
|
manifest = {
|
|
"created_utc": datetime.now(timezone.utc).isoformat(),
|
|
"sample_count": len(samples),
|
|
"group_count": len({sample.group_id for sample in samples}),
|
|
"folds": len(folds),
|
|
"seeds": args.seeds,
|
|
"device": str(device),
|
|
"gpu_name": torch.cuda.get_device_name(device) if device.type == "cuda" else None,
|
|
"python": platform.python_version(),
|
|
"torch": torch.__version__,
|
|
"feature_dir": str(args.feature_dir),
|
|
"baseline_dir": str(args.baseline_dir),
|
|
"variants": stage_manifests,
|
|
"elapsed_seconds": time.time() - start_time,
|
|
"fixed_methods_unchanged": ["M1", "M2"],
|
|
"feature_extraction_changed": False,
|
|
}
|
|
(args.output_dir / "run_manifest.json").write_text(
|
|
json.dumps(manifest, ensure_ascii=False, indent=2, allow_nan=False), encoding="utf-8"
|
|
)
|
|
print(
|
|
f"[all done] elapsed={manifest['elapsed_seconds']:.1f}s output={args.output_dir}",
|
|
flush=True,
|
|
)
|
|
return manifest
|
|
|
|
|
|
def build_parser() -> argparse.ArgumentParser:
|
|
project_dir = Path(__file__).resolve().parents[1]
|
|
parser = argparse.ArgumentParser(
|
|
description="Train staged M3/M4 alignment-loss variants and stronger temporal probes."
|
|
)
|
|
parser.add_argument("--feature-dir", type=Path, default=project_dir / "outputs/q1_features/features")
|
|
parser.add_argument("--manifest", type=Path, default=project_dir / "outputs/audit/manifest.csv")
|
|
parser.add_argument("--baseline-dir", type=Path, default=project_dir / "outputs/method_comparison")
|
|
parser.add_argument("--output-dir", type=Path, default=project_dir / "outputs/alignment_v2")
|
|
parser.add_argument("--device", default="auto", help="auto, cpu, or a torch device such as cuda:0")
|
|
parser.add_argument("--expected-samples", type=int, default=100)
|
|
parser.add_argument("--folds", type=int, default=5)
|
|
parser.add_argument("--seeds", type=int, nargs="+", default=[42, 3407, 2026])
|
|
parser.add_argument("--grid-size", type=int, default=50)
|
|
parser.add_argument("--hidden-size", type=int, default=128)
|
|
parser.add_argument("--heads", type=int, default=4)
|
|
parser.add_argument("--dropout", type=float, default=0.1)
|
|
parser.add_argument("--batch-size", type=int, default=8)
|
|
parser.add_argument("--probe-batch-size", type=int, default=16)
|
|
parser.add_argument("--epochs", type=int, default=50)
|
|
parser.add_argument("--patience", type=int, default=8)
|
|
parser.add_argument("--learning-rate", type=float, default=1e-4)
|
|
parser.add_argument("--mask-ratio", type=float, default=0.2)
|
|
parser.add_argument("--retrieval-probe-epochs", type=int, default=20)
|
|
parser.add_argument("--reconstruction-probe-epochs", type=int, default=25)
|
|
parser.add_argument("--retrieval-tolerance", type=int, default=1)
|
|
parser.add_argument("--shuffle-repeats", type=int, default=5)
|
|
parser.add_argument("--mvr-epsilon", type=float, default=0.02)
|
|
parser.add_argument("--example-id", default="-tPCytz4rww/12")
|
|
return parser
|
|
|
|
|
|
def main() -> int:
|
|
args = build_parser().parse_args()
|
|
manifest = run(args)
|
|
print(json.dumps(manifest, ensure_ascii=False, indent=2))
|
|
return 0
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|