diff options
Diffstat (limited to 'experiments')
| -rwxr-xr-x | experiments/analyze_verified.py | 36 | ||||
| -rwxr-xr-x | experiments/ep_matched_sweep.sh | 37 | ||||
| -rw-r--r-- | experiments/run.py | 5 |
3 files changed, 77 insertions, 1 deletions
diff --git a/experiments/analyze_verified.py b/experiments/analyze_verified.py index 8a1a841..50863b5 100755 --- a/experiments/analyze_verified.py +++ b/experiments/analyze_verified.py @@ -124,6 +124,40 @@ def print_local_baselines(root): f"{fmt([r['final']['wall_s'] for r in rows], digits=1)} |") +def forward_parameter_count(record): + args = record["args"] + n_in = args.get("n_in") or {"mnist": 784, "fmnist": 784, "cifar10": 3072}[args["dataset"]] + sizes = [n_in] + [args["width"]] * args["depth"] + [10] + return sum(n_out * n_in_ + n_out for n_in_, n_out in zip(sizes[:-1], sizes[1:])) + + +def best_epoch_accuracy(record): + values = [step["test_acc"] for step in record.get("steps", []) if "test_acc" in step] + values.append(record["final"]["test_acc"]) + return max(values) + + +def print_ep_match(root): + records = read_many([ + os.path.join(root, "ep_original_v2_*.json"), + os.path.join(root, "ep_match_v1_*.json"), + ]) + groups = {} + for record in records: + args = record["args"] + key = (method_name(record), args["depth"], args["width"], + forward_parameter_count(record), args["epochs"]) + groups.setdefault(key, []).append(record) + print("\n\n## EP near-parameter-matched comparison\n") + print("| method | depth × width | forward params | epochs | n | last acc (%) | best acc (%) | wall (s) |") + print("|:---|:---|---:|---:|---:|---:|---:|---:|") + for (method, depth, width, params, epochs), rows in sorted(groups.items()): + print(f"| {method} | {depth} × {width} | {params:,} | {epochs} | {len(rows)} | " + f"{fmt([r['final']['test_acc'] for r in rows], percent=True)} | " + f"{fmt([best_epoch_accuracy(r) for r in rows], percent=True)} | " + f"{fmt([r['final']['wall_s'] for r in rows], digits=1)} |") + + def provenance_audit(root): records = read_many([ os.path.join(root, "scale_v3_*.json"), @@ -132,6 +166,7 @@ def provenance_audit(root): os.path.join(root, "pepita_tuned_v1_*.json"), os.path.join(root, "baseline_budget_v1_mnist_ff_*.json"), os.path.join(root, "ep_original_v2_*.json"), + os.path.join(root, "ep_match_v1_*.json"), ]) dirty = [r["_path"] for r in records if r.get("provenance", {}).get("git_dirty") is not False] print(f"\n\nProvenance: {len(records)} files audited; {len(dirty)} dirty or unknown.") @@ -146,6 +181,7 @@ def main(): print_scaling(args.results) print_nuisance(args.results) print_local_baselines(args.results) + print_ep_match(args.results) provenance_audit(args.results) diff --git a/experiments/ep_matched_sweep.sh b/experiments/ep_matched_sweep.sh new file mode 100755 index 0000000..9c23034 --- /dev/null +++ b/experiments/ep_matched_sweep.sh @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +# Near-parameter-matched comparison against the canonical d1/w500 EP run. +# All methods see the same first 50k MNIST training examples for 25 epochs. +# Usage: ep_matched_sweep.sh <gpu> "<methods>" "<seeds>" [prefix] +set -eu + +cd "$(dirname "$0")/.." +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +GPU="${1:?GPU index required}" +METHODS="${2:-bp dfa sdil}" +SEEDS="${3:-0}" +PREFIX="${4:-ep_match_v1}" + +export CUDA_VISIBLE_DEVICES="$GPU" +export OMP_NUM_THREADS=2 +mkdir -p results logs/baselines + +for seed in $SEEDS; do + for method in $METHODS; do + tag="${PREFIX}_mnist_${method}_w256_d3_s${seed}" + result="results/${tag}.json" + log="logs/baselines/${tag}.log" + if [[ -f "$result" ]]; then + echo "skip $tag (result exists)" + continue + fi + echo ">>> $tag $(date --iso-8601=seconds) gpu=$GPU" + "$PY" experiments/run.py \ + --mode "$method" --dataset mnist --depth 3 --width 256 \ + --epochs 25 --batch_size 128 --train_examples 50000 \ + --eta 0.05 --eta_A 0.02 --eta_P 0.002 --momentum 0.9 \ + --pert_every 4 --pert_ndirs 16 --pert_mode simultaneous \ + --nuis_rho 0 --residual 0 --act tanh --seed "$seed" \ + --tag "$tag" --outdir results > "$log" 2>&1 + grep -h DONE "$log" + done +done diff --git a/experiments/run.py b/experiments/run.py index eea3f91..09e5a30 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -76,7 +76,8 @@ def train(args): device = args.device torch.manual_seed(args.seed) train_loader, test_loader, n_in, n_out = get_dataset( - args.dataset, batch_size=args.batch_size, device=device) + args.dataset, batch_size=args.batch_size, device=device, + train_limit=args.train_examples or None) args.n_in = n_in net, cfg = build(args, device) @@ -164,6 +165,8 @@ def get_args(): p.add_argument("--residual", type=int, default=0) # skip connections (deep no-BN) p.add_argument("--epochs", type=int, default=15) p.add_argument("--batch_size", type=int, default=128) + p.add_argument("--train_examples", type=int, default=0, + help="0 uses the full training split") p.add_argument("--eta", type=float, default=0.05) p.add_argument("--eta_A", type=float, default=0.02) p.add_argument("--eta_P", type=float, default=0.002) |
