summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rwxr-xr-xexperiments/analyze_verified.py36
-rwxr-xr-xexperiments/ep_matched_sweep.sh37
-rw-r--r--experiments/run.py5
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)