Files
modeling_zhaocui/deep_learning/Q1/q1/train_alignment_variants.py

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())