summaryrefslogtreecommitdiff
path: root/scripts/downstream_capacity_sweep.py
diff options
context:
space:
mode:
Diffstat (limited to 'scripts/downstream_capacity_sweep.py')
-rw-r--r--scripts/downstream_capacity_sweep.py27
1 files changed, 22 insertions, 5 deletions
diff --git a/scripts/downstream_capacity_sweep.py b/scripts/downstream_capacity_sweep.py
index 6e42550..a5915ea 100644
--- a/scripts/downstream_capacity_sweep.py
+++ b/scripts/downstream_capacity_sweep.py
@@ -37,11 +37,14 @@ class RunConfig:
optimizer: str
init_seeds: int
feedback_seeds: int
+ init_seed_offset: int
+ feedback_seed_offset: int
data_seed: int
noise_std: float
feedback_scale: str
capacity_q: float
jacobian_lambda_rel: float
+ skip_jacobian: bool
device: str
torch_threads: int
outdir: str
@@ -120,6 +123,8 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--optimizer", choices=["sgd", "adam"], default="sgd")
parser.add_argument("--init-seeds", type=int, default=3)
parser.add_argument("--feedback-seeds", type=int, default=5)
+ parser.add_argument("--init-seed-offset", type=int, default=0)
+ parser.add_argument("--feedback-seed-offset", type=int, default=0)
parser.add_argument("--data-seed", type=int, default=0)
parser.add_argument("--noise-std", type=float, default=0.0)
parser.add_argument(
@@ -129,6 +134,7 @@ def parse_args() -> argparse.Namespace:
)
parser.add_argument("--capacity-q", type=float, default=0.01)
parser.add_argument("--jacobian-lambda-rel", type=float, default=1e-3)
+ parser.add_argument("--skip-jacobian", action="store_true")
parser.add_argument("--device", choices=["cpu", "cuda"], default="cpu")
parser.add_argument("--torch-threads", type=int, default=0)
parser.add_argument(
@@ -158,11 +164,14 @@ def parse_config(args: argparse.Namespace) -> RunConfig:
optimizer=args.optimizer,
init_seeds=args.init_seeds,
feedback_seeds=args.feedback_seeds,
+ init_seed_offset=args.init_seed_offset,
+ feedback_seed_offset=args.feedback_seed_offset,
data_seed=args.data_seed,
noise_std=args.noise_std,
feedback_scale=args.feedback_scale,
capacity_q=args.capacity_q,
jacobian_lambda_rel=args.jacobian_lambda_rel,
+ skip_jacobian=args.skip_jacobian,
device=args.device,
torch_threads=args.torch_threads,
outdir=str(args.outdir),
@@ -792,7 +801,7 @@ def main() -> None:
flush=True,
)
for init_index in range(config.init_seeds):
- init_seed = 10_000 + init_index
+ init_seed = 10_000 + config.init_seed_offset + init_index
initial_weights = initialize_weights(config, width, init_seed)
bp_weights = train(
initial_weights,
@@ -805,9 +814,12 @@ def main() -> None:
)
bp_train = mse(bp_weights, x_train, y_train)
bp_test = mse(bp_weights, x_test, y_test)
- d_eff, hard_rank, lam = jacobian_effective_dimension(
- bp_weights, x_probe, config.jacobian_lambda_rel
- )
+ if config.skip_jacobian:
+ d_eff, hard_rank, lam = 0.0, 0, 0.0
+ else:
+ d_eff, hard_rank, lam = jacobian_effective_dimension(
+ bp_weights, x_probe, config.jacobian_lambda_rel
+ )
redundancy = p_count - d_eff - burden
rows.append(
RunRow(
@@ -835,7 +847,12 @@ def main() -> None:
)
for feedback_index in range(config.feedback_seeds):
- feedback_seed = 100_000 + init_index * 1000 + feedback_index
+ feedback_seed = (
+ 100_000
+ + config.feedback_seed_offset
+ + init_index * 1000
+ + feedback_index
+ )
feedback = init_feedback(config, width, feedback_seed)
fa_weights = train(
initial_weights,