建立分批同步基线(基础文件)
This commit is contained in:
@@ -0,0 +1,647 @@
|
||||
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())
|
||||
Reference in New Issue
Block a user