summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/analyze_coupled_ladder_scaling.py34
-rw-r--r--experiments/coupled_ladder_scaling.py7
2 files changed, 37 insertions, 4 deletions
diff --git a/experiments/analyze_coupled_ladder_scaling.py b/experiments/analyze_coupled_ladder_scaling.py
index 7f32c50..456b6ed 100644
--- a/experiments/analyze_coupled_ladder_scaling.py
+++ b/experiments/analyze_coupled_ladder_scaling.py
@@ -57,6 +57,11 @@ def parse_args() -> argparse.Namespace:
default=Path("results/figs/figure_clln_scaling_pilot"))
parser.add_argument("--bootstrap-replicates", type=int, default=20000)
parser.add_argument("--bootstrap-seed", type=int, default=20260829)
+ parser.add_argument(
+ "--confirmatory",
+ action="store_true",
+ help="label the output as confirmation after validating its sources",
+ )
return parser.parse_args()
@@ -240,9 +245,26 @@ def trace_summary(
def build_analysis(
- records: list[dict], *, replicates: int, seed: int
+ records: list[dict], *, replicates: int, seed: int,
+ confirmatory: bool = False,
) -> dict:
sizes = sorted({record["side"] for record in records})
+ task_size_draw_counts = {
+ len({
+ record["device_seed"] for record in records
+ if record["side"] == side
+ and record["task_index"] == task_index
+ })
+ for side in sizes
+ for task_index in {
+ record["task_index"] for record in records
+ if record["side"] == side
+ }
+ }
+ if len(task_size_draw_counts) != 1:
+ raise ValueError(
+ "every task-size cell must contain the same number of device draws")
+ component_draws_per_task_size = task_size_draw_counts.pop()
edges = np.asarray([
next(
record["learnable_edges"] for record in records
@@ -280,11 +302,12 @@ def build_analysis(
}
return {
"analysis": "digital_coupled_ladder_scaling_pilot_analysis",
- "confirmatory": False,
+ "confirmatory": confirmatory,
"bootstrap": {
"unit": "task; component draws averaged within task",
"task_clusters": len({
record["task_index"] for record in records}),
+ "component_draws_per_task_size": component_draws_per_task_size,
"replicates": replicates,
"seed": seed,
"interval": "percentile 95%",
@@ -453,10 +476,14 @@ def plot_figure(path: Path, analysis: dict) -> None:
frameon=False,
bbox_to_anchor=(0.5, 1.02),
)
+ task_count = analysis["bootstrap"]["task_clusters"]
+ component_draws = analysis["bootstrap"]["component_draws_per_task_size"]
+ evidence_label = "Confirmation" if analysis["confirmatory"] else "Exploratory pilot"
figure.text(
0.995,
0.005,
- "Exploratory pilot: 5 tasks, 1 component draw per task and size",
+ f"{evidence_label}: {task_count} tasks, {component_draws} component "
+ "draws per task and size",
ha="right",
va="bottom",
fontsize=7,
@@ -479,6 +506,7 @@ def main() -> None:
records,
replicates=args.bootstrap_replicates,
seed=args.bootstrap_seed,
+ confirmatory=args.confirmatory,
)
analysis["sources"] = [str(path) for path in source_paths]
args.output_analysis.parent.mkdir(parents=True, exist_ok=True)
diff --git a/experiments/coupled_ladder_scaling.py b/experiments/coupled_ladder_scaling.py
index 680b855..0ccee60 100644
--- a/experiments/coupled_ladder_scaling.py
+++ b/experiments/coupled_ladder_scaling.py
@@ -281,6 +281,11 @@ def parse_args() -> argparse.Namespace:
default=2.3,
)
parser.add_argument("--workers", type=int, default=8)
+ parser.add_argument(
+ "--confirmatory",
+ action="store_true",
+ help="mark a run whose protocol was frozen before endpoints were read",
+ )
return parser.parse_args()
@@ -350,7 +355,7 @@ def main() -> None:
record["side"], record["task_index"], record["device_seed"]))
report = {
"analysis": "digital_coupled_learning_size_ladder",
- "confirmatory": False,
+ "confirmatory": args.confirmatory,
"autodiff_used": False,
"source_protocol": str(args.protocol),
"protocol": {