From 1e750ea2102d76ad3dea1618655b926911f4f247 Mon Sep 17 00:00:00 2001 From: Srinivas Venkatanarayanan <17814051+Vi-Sri@users.noreply.github.com> Date: Fri, 25 Sep 2026 16:10:15 -0400 Subject: [PATCH 1/2] feat(examples): scaffold hard-case diagnostics Plan controlled interventions on known hard examples and add a paired CPU smoke runner using the public model-editing API. [force ci] --- docs/proposals/hard-example-diagnostics.md | 173 +++++++++++++ .../hard-example-diagnostics/CONTRACTS.md | 74 ++++++ .../hard-example-diagnostics/README.md | 59 +++++ .../hard-example-diagnostics/experiment.json | 42 ++++ .../hard-example-diagnostics/plan.py | 66 +++++ .../hard-example-diagnostics/smoke.py | 232 ++++++++++++++++++ .../smoke_config.json | 28 +++ 7 files changed, 674 insertions(+) create mode 100644 docs/proposals/hard-example-diagnostics.md create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/CONTRACTS.md create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/experiment.json create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke.py create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke_config.json diff --git a/docs/proposals/hard-example-diagnostics.md b/docs/proposals/hard-example-diagnostics.md new file mode 100644 index 00000000..921ec87b --- /dev/null +++ b/docs/proposals/hard-example-diagnostics.md @@ -0,0 +1,173 @@ +# Make one failure understandable — and test whether we can fix it + +**WeightsLab · meeting brief · 25 September 2026** + +**Status:** proposal and runnable scaffolding; real-data results are pending. + +**Foundation:** [merged model-editing API, PR #287](https://github.com/GrayboxTech/weightslab/pull/287). + +## The outcome we want + +An engineer selects a meaningful failure, compares it with examples the model +gets right, inspects the relevant layers, and tries a controlled intervention. +The result shows what improved, what regressed, and exactly what changed. + +Our first demonstration should answer: **Can an intervention improve one +coherent failure mode beyond simply training longer, while preserving ordinary +cases?** Improvement is a hypothesis to test, not an outcome we assume. + +The transcript's rare-tiger example is the product story: explain a recurring +failure and test a remedy. An attribution heatmap alone cannot establish that +new neurons learned a particular concept or that insufficient capacity caused +the original failure. + +## Decisions and todos for today's meeting + +Suggested 25-minute agenda. Owners below are proposed, not assigned commitments. + +| Time | Decision / todo | Proposed owner | Leave with | +|---|---|---|---| +| 0–5 min | Choose one dataset and failure group | Vi-Sri + ML reviewer | Waterbirds first; explicit group definition | +| 5–10 min | Agree the comparison and success criteria | Vi-Sri + ML reviewer | Equal training budget; ordinary-case regression bound | +| 10–15 min | Agree the first visual journey | Vi-Sri + Studio maintainer | Case selection → layer inspection → branch → comparison | +| 15–20 min | Review API and history ownership | Vi-Sri + backend maintainer | Layer identity, sampled signals, checkpoint/event contract | +| 20–25 min | Choose the next demo checkpoint | Team | GPU time, reviewer, next biweekly meeting and fallback day | + +- [ ] Confirm Waterbirds as the first controlled experiment; Oxford Pets next. +- [ ] Pick the target group using training/validation evidence, then freeze the choice. +- [ ] Confirm frozen ViT-B/16 + editable MLP head as the first editing scope. +- [ ] Agree whether the proposed +5 percentage-point target and ≤1-point common-case + regression budget are meaningful for this demo. +- [ ] Confirm where longer experiments belong and what small checks belong in CI. +- [ ] Request the existing Studio demo/reference so the prototype can follow its interaction style. +- [ ] Confirm the team member who will review model diagnostics and the Studio API contract. + +## Three experiments, in order + +| Experiment | Concrete question | Setup and intervention | Evidence / stop condition | +|---|---|---|---| +| **E1 · Waterbirds: unusual backgrounds** | Does head capacity or sample exposure explain failure on an atypical group? | Pretrained frozen ViT-B/16; fork one trained head checkpoint into continued training, wider head, group-balanced sampling, and wider head + balanced sampling. | Accuracy for every class/background group, margin, loss, corrected and newly broken examples. If the baseline has no recurring group failure, stop and report that before changing the protocol. | +| **E2 · Representation bottleneck** | Is the frozen representation limiting adaptation? | On the same selected cases, compare continued training with an unchanged head versus unfreezing the last ViT block. Keep architecture fixed and use a conventional PyTorch fine-tuning path until full-model editing is capability-tested. | Per-layer update/gradient traces and class-conditioned input attribution. Treat this as a separate representation experiment; it does not validate structural transformer edits. | +| **E3 · Oxford Pets: natural appearance variation** | Does the workflow transfer to a visually recognizable, naturally occurring subgroup? | Breed classification; manually review a coherent pose/occlusion subgroup and its confusable breed, then repeat the selected controlled comparison. | Fixed, labelled subgroup across train/validation/test. If subgroup support is too small, report an exploratory case study rather than a population-level result. | + +Waterbirds provides bird-class/background groups and a reproducible setting for +atypical-context failures. It uses **composited images**, so it is a controlled +first experiment, not evidence about natural wildlife rarity. Preserve the +published splits and account for the known label notes in the authors' repository. +Source: [Waterbirds authors' dataset and protocol](https://github.com/kohpangwei/group_DRO#waterbirds). + +Oxford-IIIT Pets provides 37 breeds, head boxes, and foreground trimaps. Pose or +occlusion subgroup labels would be our annotations, not supplied dataset labels. +Source: [Oxford-IIIT Pet dataset](https://www.robots.ox.ac.uk/~vgg/data/pets/). + +Pin `ViT_B_16_Weights.IMAGENET1K_V1` and its preprocessing rather than a moving +default. Model reference: [Torchvision ViT-B/16](https://docs.pytorch.org/vision/stable/models/generated/torchvision.models.vit_b_16.html). + +## What the engineer sees + +```mermaid +flowchart LR + A[Select known hard cases] --> B[Compare with correct examples] + B --> C[Inspect predictions and layer signals] + C --> D[Record a hypothesis] + D --> E[Fork the same checkpoint] + E --> F[Continue training: control] + E --> G[Edit model or training sample policy] + F --> H[Compare on fixed held-out cases] + G --> H + H --> I[Keep, revise, or reject intervention] + I --> J[Replayable experiment history] + classDef proposed fill:#e8f1ff,stroke:#3064ae,color:#173153; + class A,B,C,D,E,H,I,J proposed; +``` + +Blue nodes describe the new diagnostics and comparison workflow. Model edits, +the ledger, and step-level model signals already have backend foundations. + +The first screen should have four linked areas: + +1. **Cases:** original image, label, prediction, confidence, margin, group, + and human notes. A hard positive is a same-class example distant from an + anchor; a hard negative is a different-class example close to it under a + stated representation. Store the anchor and distance definition. +2. **Model inspector:** graph with selected layer, dimensions, frozen state, + activation summaries, gradient norms, and weight-change history. Curves + suggest hypotheses; they do not automatically diagnose capacity or label errors. +3. **Intervention:** checkpoint, target layer, edit preview, sample policy, + and engineer's reason. The engineer chooses whether a valid difficult + sample stays, is emphasized, or is excluded from a training branch. +4. **Comparison and history:** baseline/control/intervention side by side; + subgroup results, ordinary-case regressions, selected images and attributions, + plus a timeline of the checkpoint and every edit. + +Start from a small known case set. Automated embedding-based discovery follows +after we establish which comparisons actually help engineers make decisions. + +## Make the result defensible + +- Select the failure mode on train/validation; lock the test manifest before + intervention selection. Do not remove evaluation cases because they are hard. +- Fork identical model weights. Match seeds, update counts, batch schedule, + preprocessing, and optimizer policy between paired arms. Log changed sampling + separately from changed capacity. Reinitialize optimizers equally when a + shape-changing edit cannot preserve optimizer state. +- Compare **post-training edit versus post-training control**, not only edited + versus pre-training checkpoint. Also evaluate immediately after editing to + separate the edit's instantaneous effect from subsequent learning. +- Run three paired seeds initially. Show group counts, per-seed results, and + uncertainty. Three seeds and a small subgroup are preliminary evidence. +- Proposed pilot target: ≥5 percentage points on the selected subgroup versus + continued training, ≤1 point common-case regression, and consistent direction + across seeds. Agree these thresholds before looking at final test results; + meeting them alone does not establish statistical significance. +- Show both corrections and regressions. A null result is useful: record whether + more ordinary training, balanced exposure, or changing representation helped. + +**Attribution boundary:** a frozen, deterministic backbone produces the same +attention for the same input after a head-only edit. Class-conditioned input +gradients may change because the head changed. Compute attribution through the +full image→backbone→head graph, with input gradients enabled; cached embeddings +alone cannot yield pixel attribution. Attention is a view of token interactions, +not a literal picture of everything the model "sees". Pin target class, baseline, +preprocessing and color scale across comparisons; record approximation error +for [Integrated Gradients](https://captum.ai/docs/extension/integrated_gradients). + +If widening helps, a later ablation can disable only the newly added units to +test their contribution. Even that supports a contribution claim, not a claim +that an individual neuron represents "stripes". + +## Delivery plan + +These are work packages, not promised calendar dates. Start the first GPU pilot +after today's dataset and scope decisions. + +| Package | Deliverable | Acceptance / handoff | +|---|---|---| +| **Now · planning draft** | This brief, experiment matrix generator, local synthetic smoke runner, proposed data contracts | Runnable scaffolding; synthetic results clearly labelled | +| **1 · baseline and cases** | Pretrained ViT checkpoint, dataset/split hashes, selected validation cases, locked test set | Repeatable subgroup failure; enough examples to evaluate it | +| **2 · controlled interventions** | Four E1 arms × three seeds, immediate-edit and post-training results | Checkpoint parity, optimizer rebinding, every-group metrics and regressions | +| **3 · diagnostic prototype** | Image comparison, layer histories, attribution, edit/history timeline | Same case and layer refer to the same snapshot across views | +| **4 · Studio integration** | Versioned requests, bounded signal fetches, edit acknowledgements and refreshed graph | Browser/server training pause, edit, resume and synchronization verified end to end | +| **Later · discovery** | Candidate hard-pair mining using the task model; optional external embeddings | Discovery quality evaluated separately; human review remains available | + +Vi-Sri's proposed scope is the experiment, model representation, API/data +contracts and diagnostic prototype. Backend and Studio maintainers review how +those contracts join their existing systems. Keep each package reviewable in +its own PR; agree test placement with maintainers before adding long GPU jobs. + +## Five-minute demonstration script + +1. Show several examples of the same validation failure mode and a correct reference. +2. Point to the prediction margin and one informative layer comparison. +3. State the hypothesis, show the common checkpoint, and preview the intervention. +4. Replay recorded control/intervention runs; live training is optional. +5. Show held-out improvements **and** regressions, then the history needed to reproduce them. + +Opening line for today's meeting: + +> The editing API is merged. Next I want to make one recurring failure +> inspectable and test whether editing actually helps beyond training longer. +> We can start with known hard examples, keep the engineer in control, and build +> the API and visual comparison needed to bring that workflow into Studio. + +Implementation entry point: [diagnostics scaffold](../../weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md). diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/CONTRACTS.md b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/CONTRACTS.md new file mode 100644 index 00000000..2671f13b --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/CONTRACTS.md @@ -0,0 +1,74 @@ +# Proposed diagnostic contracts — version 0 + +Design proposal, not a public API commitment. Prototype with JSON artifacts; +agree the schema before changing the shared proto and regenerating both clients. + +## Objects and ownership + +| Object | Required identity / fields | Purpose | +|---|---|---| +| CaseSet | dataset/version, split, manifest SHA-256, case IDs, subgroup definition, selection checkpoint, reviewer note | Reproduce the same cases; detect split leakage | +| Case | stable sample ID, label, group, media reference, optional anchor ID and pair type | Keep hard positives/negatives meaningful relative to a representation | +| ModelSnapshot | run ID, checkpoint SHA-256, architecture revision, model age, backbone/preprocessing fingerprints, graph | Tie measurements to exact weights and architecture | +| LayerRef | snapshot ID, module path, live layer ID, shape | Live IDs are valid only within a wrapped model; resolve module paths again after reload | +| CaseDiagnostic | snapshot, case ID, logits, loss, true-label margin, prediction, selected-layer summaries | Compare the same example before and after | +| LayerHistory | snapshot, layer ref, step, scope, sample/batch IDs where applicable, gradient/activation/update values | Distinguish per-case observations from batch/step aggregates | +| Intervention | event ID, parent checkpoint, expected architecture revision, target path/ID, operation, arguments, reason, before/after shapes, optimizer policy, status | Explain and replay the engineer's action | +| Comparison | control/intervention run IDs, shared parent checkpoint, split hash, training budgets, per-group counts/metrics, corrected/regressed IDs | Separate intervention effects from ordinary continued training | + +Missing/unavailable measurements must be `null` with an availability reason, +not zero. Version artifacts, keep numeric values finite, and use UTC timestamps. + +## Reuse today's backend + +- Structure: `model.get_model_graph()` and `model.get_layer_info()`. +- Per-sample evidence: `wl.save_signals` and sample history queries. +- Per-step health: `wl.watch_or_edit(..., track_model_signals=True)` or + `wl.track_model_signals`; never broadcast a batch gradient norm onto samples. +- Editable operations: public model editing methods from PR #287. +- Lifecycle: `wl.guard_training_context`, `wl.guard_testing_context`, `wl.start_training`. + +Extra work: checkpoint-to-checkpoint weight deltas, bounded per-case activation +capture, attribution metadata, experiment branch comparisons and event history. +For resized tensors, compare retained rows/columns and summarize new parameters +separately; never subtract different shapes or present new weights as drift. + +## Proposed UI requests (conceptual, not existing endpoint names) + +1. `ListCases(case_set, cursor, limit)` returns a bounded case page. +2. `InspectCases(snapshot, case_ids, layer_paths, requested_signals)` returns a + bounded diagnostic payload; cache by checkpoint, preprocessing and case IDs. +3. `GetLayerHistory(run, layer_path, step_range, resolution)` downsamples curves. +4. `PreviewIntervention(snapshot, operation)` returns affected shapes and support status. +5. `ApplyIntervention(expected_revision, operation, reason)` pauses at a training + boundary and applies under the existing architecture lock. Reject stale + revisions and unsupported operations before mutation. On failure, report + the error and restore a known checkpoint before allowing training to resume. +6. `CompareRuns(control, intervention, case_set)` joins results by stable case ID. + +An edit acknowledgement must carry the new architecture revision, affected +layers, optimizer rebinding status and refreshed graph. The browser must discard +stale diagnostics and refetch that revision before enabling another edit. Save +an event only as successful after shape checks and a forward pass succeed. + +## Attribution contract + +Store method, target class/logit, baseline definition, input preprocessing, +checkpoint, seed, numerical approximation settings and shared display scale. +For attention include layer/head/token selection. For Integrated Gradients +include convergence delta. A CLS embedding is a vector, not a spatial heatmap. + +Frozen ViT + head edit: backbone attention and backbone feature outputs should +stay fixed for an identical input in eval mode. Output-conditioned input +attribution can change. Verify this distinction in the UI and experiment tests. +If perturbation is added, label whether the changed object is a Q/K weight, +an attention logit/probability, or an activation; do not call them interchangeable. + +## Integration acceptance + +- Select the same case and layer in UI/backend; snapshot and revision agree. +- Pause, edit, validate propagation and optimizer parameter references, resume. +- Refetch graph and diagnostics; ignore stale replies from before the edit. +- Recover the same predictions from a saved checkpoint plus edit/event recipe. +- Record failures without claiming a completed intervention or improvement. +- Compare fixed held-out cases and display corrected as well as regressed samples. diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md new file mode 100644 index 00000000..aaacb46b --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md @@ -0,0 +1,59 @@ +# Hard-example diagnostics + +Planning draft for the project after the model-editing API. Read the +[meeting brief](../../../../../docs/proposals/hard-example-diagnostics.md) first. + +## What runs today + +- `plan.py`: expands `experiment.json` into 12 planned Waterbirds runs; standard + library only. It does not download data or launch training. +- `smoke.py`: runs two actual CPU training branches on a tiny synthetic fixture, + starting from identical weights. One continues unchanged; the other uses the + public WeightsLab neuron-addition API. Exports predictions, graph snapshots, + checkpoint identity, edit history and ordinary/rare-group metrics. +- `CONTRACTS.md`: proposed case, diagnostic, intervention and comparison contracts + for the backend and Studio. No new RPC or Studio screen is implemented yet. + +Synthetic fixture metrics test the plumbing. They are **not Waterbirds, ViT, +or evidence that editing improves real rare cases**. No image attribution or +data download is implemented in this scaffold. + +From the repository root, with WeightsLab's dependencies installed: + +```bash +python weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py \ + --output /tmp/weightslab-hard-example-plan.json + +WL_NO_TELEMETRY=1 PYTHONPATH=. python \ + weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke.py \ + --output-dir /tmp/weightslab-hard-example-smoke +``` + +Use a fresh output directory for each smoke run. Output is `report.json` and +`checkpoint.pt` alongside WeightsLab's state. The runner fails if propagation, +optimizer parameter binding, checkpoint parity, or finite training checks fail. +It does not assert that the widened model wins. + +The smoke fork restores learned head parameters, starts each branch's local step +count at zero, and uses fresh SGD without momentum in both arms. It is not a full +WeightsLab checkpoint/replay implementation; the report records parent model age +separately. Signal tracking starts after an edit so hooks bind to current tensors. + +## Next implementation steps + +- [ ] Add a Waterbirds adapter with versioned metadata and unchanged official splits. +- [ ] Cache features from a pinned pretrained ViT-B/16 in evaluation mode; retain + image IDs, image paths, preprocessing and backbone weight fingerprint. +- [ ] Select a known failure group on validation and freeze the case manifest. +- [ ] Extend the paired runner to the four planned sampling/capacity arms. +- [ ] Record per-case loss/margin and selected-layer activations; reuse + `wl.save_signals` and `wl.track_model_signals` for their appropriate scopes. +- [ ] Capture step histories before and after edits; verify signal hooks rebind + to replaced parameters and do not duplicate events. +- [ ] Add class-conditioned image attribution through the full frozen backbone. +- [ ] Add a read-only comparison view and intervention history, then Studio transport. +- [ ] Re-run the full ViT capability probe before any structural backbone editing. + +Prototype lives under `wl-model-editing` so later architectures can share the +contracts. Keep real GPU experiment reports outside source control; commit the +recipe, split manifest hashes and compact reviewed summaries instead. diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/experiment.json b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/experiment.json new file mode 100644 index 00000000..825ca451 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/experiment.json @@ -0,0 +1,42 @@ +{ + "schema_version": 1, + "experiment_name": "waterbirds_known_failure", + "status": "proposal_awaiting_dataset_and_baseline", + "dataset": { + "name": "waterbirds", + "version": "waterbird_complete95_forest2water2", + "source": "https://github.com/kohpangwei/group_DRO#waterbirds", + "split_policy": "preserve_official_train_validation_test", + "target_group": null, + "case_manifest_sha256": null + }, + "backbone": { + "architecture": "vit_b_16", + "weights": "IMAGENET1K_V1", + "frozen": true, + "feature_width": 768, + "preprocessing": "weights_enum.transforms()" + }, + "head": {"hidden_width": 64, "add_neurons": 16, "classes": 2}, + "device": "auto", + "root_log_dir": null, + "seeds": [17, 29, 43], + "baseline_training_steps": 500, + "training_steps_to_do": 250, + "eval_full_to_train_steps_ratio": 50, + "experiment_dump_to_train_steps_ratio": 250, + "optimizer": {"name": "SGD", "lr": 0.01, "momentum": 0.0}, + "data": {"train_loader": {"batch_size": 64}, "test_loader": {"batch_size": 128}}, + "optimizer_at_fork": "fresh_identical_optimizer_in_all_arms", + "arms": [ + {"name": "continue", "add_neurons": 0, "sampling": "original"}, + {"name": "widen", "add_neurons": 16, "sampling": "original"}, + {"name": "resample", "add_neurons": 0, "sampling": "group_balanced"}, + {"name": "widen_resample", "add_neurons": 16, "sampling": "group_balanced"} + ], + "pilot_targets_proposed": { + "target_group_gain_percentage_points": 5, + "max_common_group_regression_percentage_points": 1, + "requires_review_before_test_evaluation": true + } +} diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py new file mode 100644 index 00000000..b0d90ba1 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py @@ -0,0 +1,66 @@ +"""Expand the proposal into a run matrix; never downloads data or trains.""" + +from __future__ import annotations + +import argparse +import hashlib +import json +from pathlib import Path + + +def build_plan(config: dict) -> dict: + """Produce paired jobs with a shared baseline identifier for each seed.""" + if config["schema_version"] != 1: + raise ValueError("Unsupported experiment schema version") + seeds = config["seeds"] + arms = config["arms"] + if not seeds or len(set(seeds)) != len(seeds): + raise ValueError("Seeds must be nonempty and unique") + names = [arm["name"] for arm in arms] + if len(set(names)) != len(names) or "continue" not in names: + raise ValueError("Unique arms including the continued-training control are required") + if config["training_steps_to_do"] <= 0: + raise ValueError("Each arm needs a positive, matched training budget") + fingerprint = hashlib.sha256( + json.dumps(config, sort_keys=True).encode() + ).hexdigest() + jobs = [] + for seed in seeds: + for arm in arms: + jobs.append({ + "run_id": f"{config['experiment_name']}-{fingerprint[:8]}-{seed}-{arm['name']}", + "seed": seed, + "parent_checkpoint_id": f"baseline-{fingerprint[:8]}-seed-{seed}", + "parent_checkpoint_sha256": None, + "case_manifest_sha256": config["dataset"]["case_manifest_sha256"], + "training_steps_to_do": config["training_steps_to_do"], + "optimizer_at_fork": config["optimizer_at_fork"], + "intervention": arm, + "status": "planned_not_executed", + }) + return { + "schema_version": 1, + "artifact_type": "experiment_plan", + "config_sha256": fingerprint, + "config": config, + "jobs": jobs, + "next_gate": "Prepare dataset, validate baseline failure and lock case manifest", + } + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", type=Path, default=Path(__file__).with_name("experiment.json")) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + plan = build_plan(json.loads(args.config.read_text())) + args.output.parent.mkdir(parents=True, exist_ok=True) + # Exclusive creation prevents silently replacing an already-reviewed plan. + with args.output.open("x") as stream: + json.dump(plan, stream, indent=2, allow_nan=False) + stream.write("\n") + print(f"Planned {len(plan['jobs'])} runs (not executed): {args.output}") + + +if __name__ == "__main__": + main() diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke.py new file mode 100644 index 00000000..ec80ae6d --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke.py @@ -0,0 +1,232 @@ +"""Exercise a paired WeightsLab edit workflow on synthetic vectors, not ViT.""" + +from __future__ import annotations + +import argparse +import copy +import hashlib +import json +import math +from datetime import datetime, timezone +from pathlib import Path + +import torch +from torch import nn + +import weightslab as wl + + +def fixture(hp, split: str): + generator = torch.Generator().manual_seed(int(hp[f"{split}_data_seed"])) + count = int(hp[f"{split}_samples"]) + labels = torch.arange(count) % int(hp["num_classes"]) + rare = torch.arange(count) < round(count * float(hp[f"{split}_rare_fraction"])) + features = torch.randn(count, int(hp["feature_width"]), generator=generator) * float(hp["feature_noise"]) + sign = labels.float() * 2 - 1 + features[:, 0] += sign + features[:, 1] += sign * torch.where(rare, -1, 1) * float(hp["spurious_strength"]) + return features.to(hp["device"]), labels.to(hp["device"]), rare.to(hp["device"]) + + +def register(config: dict, output: Path, phase: str, checkpoint=None): + wl.clear_all() + phase_config = copy.deepcopy(config) + phase_config["root_log_dir"] = str(output / phase) + hp = wl.watch_or_edit(phase_config, flag="hyperparameters", defaults=phase_config) + torch.set_num_threads(int(hp["num_threads"])) + torch.manual_seed(int(hp["seed"])) + raw = nn.Sequential( + nn.Linear(int(hp["feature_width"]), int(hp["hidden_width"])), + nn.LeakyReLU(negative_slope=float(hp["negative_slope"])), + nn.Linear(int(hp["hidden_width"]), int(hp["num_classes"])), + ) + raw.task_type = "classification" + if checkpoint is not None: + raw[0].load_state_dict(checkpoint["hidden"]) + raw[2].load_state_dict(checkpoint["classifier"]) + model = wl.watch_or_edit( + raw, flag="model", device=hp["device"], + dummy_input=torch.zeros(1, int(hp["feature_width"]), device=hp["device"]), + compute_dependencies=True, forced_model_wrapping=True, + skip_previous_auto_load=True, + ) + optimizer = wl.watch_or_edit( + torch.optim.SGD(model.parameters(), lr=float(hp["optimizer"]["lr"])), flag="optimizer" + ) + wl.start_training() + return hp, model, optimizer + + +def linear_layers(model): + layers = [layer for layer in model.get_model_graph()["layers"] if layer["type"] == "Linear"] + if len(layers) != 2: + raise AssertionError("Expected exactly two Linear layers in the smoke head") + hidden = next(layer for layer in layers if layer["name"] == "0") + classifier = next(layer for layer in layers if layer["name"] == "2") + return hidden, classifier + + +def evaluate(model, data, batch_size: int) -> dict: + features, labels, rare = data + model.eval() + chunks = [] + with wl.guard_testing_context, torch.no_grad(): + for start in range(0, len(features), batch_size): + chunks.append(model(features[start:start + batch_size]).detach().cpu()) + logits = torch.cat(chunks) + if len(logits) != len(labels) or not torch.isfinite(logits).all(): + raise AssertionError("Missing or non-finite evaluation predictions") + labels, rare = labels.cpu(), rare.cpu() + predictions = logits.argmax(dim=1) + losses = nn.functional.cross_entropy(logits, labels, reduction="none") + probabilities = logits.softmax(dim=1) + other = logits.clone() + other.scatter_(1, labels[:, None], float("-inf")) + margins = logits.gather(1, labels[:, None]).squeeze(1) - other.max(dim=1).values + metrics = {} + for name, mask in (("ordinary", ~rare), ("rare", rare)): + count = int(mask.sum()) + metrics[name] = { + "count": count, + "accuracy": float((predictions[mask] == labels[mask]).float().mean()) if count else None, + "loss": float(losses[mask].mean()) if count else None, + } + return { + "model_age": model.get_age(), + "metrics": metrics, + "cases": [ + {"sample_id": f"synthetic-test-{index}", "label": int(labels[index]), + "group": "rare" if bool(rare[index]) else "ordinary", + "prediction": int(predictions[index]), "logits": logits[index].tolist(), + "confidence": float(probabilities[index].max()), "loss": float(losses[index]), + "true_label_margin": float(margins[index])} + for index in range(len(labels)) + ], + } + + +def train(model, optimizer, data, hp, steps: int) -> list[dict]: + features, labels, _ = data + model.train() + initial_age = model.get_age() + history = [] + # Tracker is installed after any edit so its hooks address the new tensors. + tracker = wl.track_model_signals(model, every_n_steps=int(hp["model_signals_every_n_steps"])) + try: + while model.get_age() - initial_age < steps: + offset = model.get_age() - initial_age + batch_size = int(hp["data"]["train_loader"]["batch_size"]) + indices = (torch.arange(batch_size, device=features.device) + offset * batch_size) % len(features) + complete = False + with wl.guard_training_context: + optimizer.zero_grad(set_to_none=True) + loss = nn.functional.cross_entropy(model(features[indices]), labels[indices]) + loss.backward() + optimizer.step() + complete = True + # The training guard may suppress errors: a skipped step is a failure. + if not complete or not math.isfinite(float(loss.detach())): + raise AssertionError("Training step failed or produced non-finite loss") + history.append({"model_age": model.get_age(), "loss": float(loss.detach())}) + finally: + tracker.remove() + return history + + +def run(config: dict, output: Path) -> dict: + hp, model, optimizer = register(config, output, "baseline") + train_data, test_data = fixture(hp, "train"), fixture(hp, "test") + train(model, optimizer, train_data, hp, int(hp["baseline_training_steps"])) + baseline = evaluate(model, test_data, int(hp["data"]["test_loader"]["batch_size"])) + hidden, classifier = linear_layers(model) + checkpoint = { + # Fork learned parameters only; WL's per-dataset tracker buffers are + # runtime instrumentation and do not belong in a fresh nn.Linear. + "hidden": {name: value.detach().cpu().clone() for name, value in + model.get_layer_by_id(hidden["id"]).named_parameters(recurse=False)}, + "classifier": {name: value.detach().cpu().clone() for name, value in + model.get_layer_by_id(classifier["id"]).named_parameters(recurse=False)}, + "parent_model_age": model.get_age(), + } + checkpoint_path = output / "checkpoint.pt" + torch.save(checkpoint, checkpoint_path) + checkpoint_hash = hashlib.sha256(checkpoint_path.read_bytes()).hexdigest() + runs = {} + for arm in ("continue", "widen"): + restored = torch.load(checkpoint_path, map_location="cpu", weights_only=True) + hp, model, optimizer = register(config, output, arm, restored) + test_batch_size = int(hp["data"]["test_loader"]["batch_size"]) + before = evaluate(model, test_data, test_batch_size) + if before["cases"] != baseline["cases"]: + raise AssertionError("Branches must start from identical checkpoint predictions") + graph_before = model.get_model_graph() + hidden, classifier = linear_layers(model) + events = [] + if arm == "widen": + model.add_neurons(hidden["id"], count=int(hp["add_neurons"])) + hidden_after = model.get_layer_info(hidden["id"], include_neurons=False) + classifier_after = model.get_layer_info(classifier["id"], include_neurons=False) + expected_width = int(hp["hidden_width"]) + int(hp["add_neurons"]) + if hidden_after["output_neurons"] != expected_width or classifier_after["input_neurons"] != expected_width: + raise AssertionError("Head widening did not propagate to the classifier") + events.append({ + "operation": "add_neurons", "layer_path": hidden["name"], + "live_layer_id": hidden["id"], "count": int(hp["add_neurons"]), + "before": hidden, "after": hidden_after, + "optimizer_policy": "fresh_SGD_no_momentum_in_both_arms", + "reason": "Synthetic plumbing check; no performance hypothesis tested", + }) + optimizer_ids = {id(p) for group in optimizer.param_groups for p in group["params"]} + if optimizer_ids != {id(p) for p in model.parameters()}: + raise AssertionError("Optimizer points at stale or missing model parameters") + immediate = evaluate(model, test_data, test_batch_size) + history = train(model, optimizer, train_data, hp, int(hp["training_steps_to_do"])) + after = evaluate(model, test_data, test_batch_size) + runs[arm] = { + "parent_checkpoint_sha256": checkpoint_hash, + "parent_model_age": checkpoint["parent_model_age"], + "branch_step_origin": 0, + "before": before, "immediately_after_edit": immediate, "after": after, + "graph_before": graph_before, "graph_after": model.get_model_graph(), + "interventions": events, "training_history": history, + "optimizer_binding_checked": True, + } + control_cases = {case["sample_id"]: case for case in runs["continue"]["after"]["cases"]} + corrected, regressed = [], [] + for case in runs["widen"]["after"]["cases"]: + control = control_cases[case["sample_id"]] + was_correct = control["prediction"] == control["label"] + is_correct = case["prediction"] == case["label"] + if is_correct and not was_correct: + corrected.append(case["sample_id"]) + elif was_correct and not is_correct: + regressed.append(case["sample_id"]) + return { + "schema_version": 0, "artifact_type": "synthetic_smoke_report", + "status": "plumbing_passed", "created_at": datetime.now(timezone.utc).isoformat(), + "claim": "Synthetic CPU workflow only; no ViT, real-data, or improvement claim", + "config": config, "torch_version": torch.__version__, + "checkpoint_sha256": checkpoint_hash, "baseline": baseline, "runs": runs, + "comparison": {"corrected_ids": corrected, "regressed_ids": regressed}, + } + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", type=Path, default=Path(__file__).with_name("smoke_config.json")) + parser.add_argument("--output-dir", type=Path, required=True) + args = parser.parse_args() + config = json.loads(args.config.read_text()) + if config["num_classes"] != 2 or config["feature_width"] < 2: + raise ValueError("The synthetic fixture requires two classes and at least two features") + args.output_dir.mkdir(parents=True, exist_ok=False) + try: + report = run(config, args.output_dir) + (args.output_dir / "report.json").write_text(json.dumps(report, indent=2, allow_nan=False) + "\n") + print(f"Synthetic workflow passed: {args.output_dir / 'report.json'}") + finally: + wl.clear_all() + + +if __name__ == "__main__": + main() diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke_config.json b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke_config.json new file mode 100644 index 00000000..6de61907 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/smoke_config.json @@ -0,0 +1,28 @@ +{ + "experiment_name": "hard_examples_synthetic_smoke", + "device": "cpu", + "root_log_dir": null, + "seed": 17, + "train_data_seed": 101, + "test_data_seed": 202, + "num_threads": 1, + "feature_width": 4, + "hidden_width": 8, + "num_classes": 2, + "add_neurons": 2, + "negative_slope": 0.1, + "train_samples": 64, + "test_samples": 40, + "train_rare_fraction": 0.125, + "test_rare_fraction": 0.5, + "feature_noise": 0.3, + "spurious_strength": 2.0, + "baseline_training_steps": 12, + "training_steps_to_do": 8, + "eval_full_to_train_steps_ratio": 0, + "experiment_dump_to_train_steps_ratio": 0, + "model_signals_every_n_steps": 1, + "skip_checkpoint_load": true, + "optimizer": {"lr": 0.05}, + "data": {"train_loader": {"batch_size": 8}, "test_loader": {"batch_size": 40}} +} From 4a9abda7fdf9f70adf96ab63199ff2e2a6358f28 Mon Sep 17 00:00:00 2001 From: Srinivas Venkatanarayanan <17814051+Vi-Sri@users.noreply.github.com> Date: Thu, 8 Oct 2026 10:04:26 -0400 Subject: [PATCH 2/2] feat(examples): guard paired experiment plans Add meeting decisions and checks for valid controls, seeds and budgets. [force ci] --- docs/proposals/hard-example-diagnostics.md | 5 +- .../hard-example-meeting-2026-10-08.md | 125 ++++++++++++++++++ .../hard-example-diagnostics/README.md | 8 +- .../hard-example-diagnostics/plan.py | 44 ++++-- .../hard-example-diagnostics/test_plan.py | 88 ++++++++++++ 5 files changed, 258 insertions(+), 12 deletions(-) create mode 100644 docs/proposals/hard-example-meeting-2026-10-08.md create mode 100644 weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_plan.py diff --git a/docs/proposals/hard-example-diagnostics.md b/docs/proposals/hard-example-diagnostics.md index 921ec87b..a39a637d 100644 --- a/docs/proposals/hard-example-diagnostics.md +++ b/docs/proposals/hard-example-diagnostics.md @@ -1,9 +1,12 @@ # Make one failure understandable — and test whether we can fix it -**WeightsLab · meeting brief · 25 September 2026** +**WeightsLab · project brief · updated 8 October 2026** **Status:** proposal and runnable scaffolding; real-data results are pending. +For the short presentation and decision checklist, open the +[8 October meeting handout](hard-example-meeting-2026-10-08.md). + **Foundation:** [merged model-editing API, PR #287](https://github.com/GrayboxTech/weightslab/pull/287). ## The outcome we want diff --git a/docs/proposals/hard-example-meeting-2026-10-08.md b/docs/proposals/hard-example-meeting-2026-10-08.md new file mode 100644 index 00000000..1c3cf0c9 --- /dev/null +++ b/docs/proposals/hard-example-meeting-2026-10-08.md @@ -0,0 +1,125 @@ +# From model editing to explainable experiments + +WeightsLab · 8 October 2026 · meeting handout + +**Decision requested:** agree one controlled failure-mode experiment and the +smallest useful inspection/comparison workflow. This is a proposal, not a +report of measured improvement. + +> The editing API is merged. Next we want to show a real failure, explain what +> we suspect, try a model or data change, and see whether it helped more than +> simply training longer. The engineer should be able to inspect both the +> improvements and the damage, and replay exactly what changed. + +## Today's five decisions + +- [ ] **First dataset:** approve Waterbirds as a controlled starting point. + Keep the published splits; identify the failure group on validation, not test. +- [ ] **Editing scope:** frozen pretrained ViT-B/16 plus an editable MLP head. + Structural changes inside attention blocks are a separate capability gate. +- [ ] **Fair comparison:** one baseline checkpoint per seed, four arms, equal + continuation budgets, and an identical optimizer-reset policy in every arm. +- [ ] **Demo contract:** show cases, layer evidence, the proposed intervention, + and before/control/after comparison with history. Review the data contract + before extending the shared proto. +- [ ] **Ownership and next checkpoint:** confirm an ML reviewer and a + backend/Studio reviewer, GPU availability, and the next demo date. Suggested + owner of the experiment, representation and prototype: Vi-Sri. + +## First experiment: capacity or sample exposure? + +Hypothesis: a recurring atypical-background failure may respond to more head +capacity, more exposure to underrepresented groups, both, or neither. Do not +assume a failure proves a capacity bottleneck. + +| Arm | Hidden head width | Training sampling | What it isolates | +|---|---|---|---| +| A · Continue | 64 | Original | Additional training alone | +| B · Widen | 80 (+16) | Original | Added capacity versus A | +| C · Resample | 64 | Group-balanced | Changed sample exposure versus A | +| D · Both | 80 (+16) | Group-balanced | Capacity versus C; sampling versus B | + +Proposed recipe: `ViT_B_16_Weights.IMAGENET1K_V1`, seeds **17, 29, 43**, +**500 baseline steps**, then **250 steps per arm**. Cache frozen features for +fast head experiments. These budgets are pilot settings, not a claim that the +baseline has converged. Check validation learning curves before locking the +protocol; record any revision before final test evaluation. + +**Evidence required:** identical starting predictions; immediate-edit versus +post-training effect; every group's accuracy and count; worst-group accuracy; +corrected and newly broken examples; per-seed paired differences; checkpoint, +split and configuration hashes. Report empirical test accuracy separately from +the training-frequency-weighted benchmark average. + +**Proposed demo target, to agree today:** at least +5 percentage points on the +selected group versus A, no more than 1 point common-case regression, and a +consistent direction across the three seeds. Define the common groups before +evaluation. These are practical pilot criteria, not statistical significance. +A null result is a valid outcome; do not keep changing the protocol until an +edit wins. + +Waterbirds deliberately combines bird foregrounds with backgrounds. That makes +the groups reproducible, but it is not a natural rare-animal benchmark. +[Dataset and evaluation protocol](https://github.com/kohpangwei/group_DRO#waterbirds). + +## What the demo should show + +```mermaid +flowchart LR + A[Surface known hard cases] --> B[Inspect predictions and layers] + B --> C[Record hypothesis and fork checkpoint] + C --> D[Continue: control] + C --> E[Edit model or sampling] + D --> F[Compare held-out cases] + E --> F + F --> G[Keep or reject with explainable history] +``` + +One selected image stays selected across four panels: + +1. **Cases:** image, label, prediction, margin, group and human notes; also a + correct reference case. Begin with known cases, not automated mining. +2. **Model:** layer path/identity, shape, activation summary, gradient norm and + weight updates. Treat these as evidence for a hypothesis, not its proof. +3. **Decision:** show affected layers, sampling choice, parent checkpoint, + optimizer policy and the engineer's reason before applying the edit. +4. **Comparison:** control and intervention results, corrections, regressions, + and a replayable event timeline. Use a read-only recorded-run prototype first. + +**Important visual boundary:** head-only edits cannot change attention in a +frozen deterministic backbone. Class-conditioned input attribution can change; +compute it through image → backbone → head with a fixed target class and +display scale. Do not draw a spatial heatmap from a cached CLS vector or call +attention a complete explanation of what the model sees. + +## Work after the meeting, in order + +| Gate | Concrete deliverable | Done when | +|---|---|---| +| 1 · Establish failure | Dataset adapter, pinned feature cache, baseline and validation case manifest | Failure is coherent and reproducible; test selection is locked | +| 2 · Test intervention | Four arms × three seeds with immediate and final snapshots | Same-parent checks, shape propagation, optimizer references and matched budgets pass | +| 3 · Make it inspectable | Case comparison, bounded layer diagnostics, attribution and history | Every view identifies the same case and model snapshot; missing data is explicit | +| 4 · Integrate Studio | Versioned requests and edit acknowledgement with architecture revision | Pause/edit/resume works; stale responses are rejected and the graph refreshes | + +The architectural contribution is the **model representation and diagnostic +API/data contract**: stable case/snapshot/layer references, bounded measurement +requests, and explainable intervention history. The prototype UI tests whether +that information helps an engineer; it can later be ported into Studio. + +Follow-ups: **E2**, keep the architecture fixed and unfreeze the last ViT block +to test a representation bottleneck; **E3**, repeat the workflow on a reviewed +natural-variation subgroup in Oxford Pets. Neither belongs in the first demo's +completion claim. [Full proposal](hard-example-diagnostics.md). + +## What is available versus planned + +- **Available:** a 12-job plan generator, proposed contracts, and a CPU + synthetic smoke runner exercising public neuron addition, dependency + propagation, optimizer binding and same-checkpoint comparisons. +- **Planned:** Waterbirds execution, pretrained feature extraction, image + attribution, persistent replay, new RPCs and the Studio comparison screen. +- **Review boundary:** draft PR [#307](https://github.com/GrayboxTech/weightslab/pull/307) + remains a planning/scaffolding PR. It does not claim to close issue #267. + +Start with the [runnable scaffold](../../weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md) +and its [proposed contracts](../../weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/CONTRACTS.md). diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md index aaacb46b..169867d3 100644 --- a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md @@ -2,11 +2,15 @@ Planning draft for the project after the model-editing API. Read the [meeting brief](../../../../../docs/proposals/hard-example-diagnostics.md) first. +For presentation, use the shorter +[meeting handout](../../../../../docs/proposals/hard-example-meeting-2026-10-08.md). ## What runs today - `plan.py`: expands `experiment.json` into 12 planned Waterbirds runs; standard - library only. It does not download data or launch training. + library only. Rejects an edited control, duplicate seeds/arms, invalid budgets + and mismatched optimizer policy. It does not download data or launch training. +- `test_plan.py`: standard-library tests for the paired plan and its guardrails. - `smoke.py`: runs two actual CPU training branches on a tiny synthetic fixture, starting from identical weights. One continues unchanged; the other uses the public WeightsLab neuron-addition API. Exports predictions, graph snapshots, @@ -21,6 +25,8 @@ data download is implemented in this scaffold. From the repository root, with WeightsLab's dependencies installed: ```bash +python weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_plan.py + python weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py \ --output /tmp/weightslab-hard-example-plan.json diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py index b0d90ba1..8c23f3d6 100644 --- a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py @@ -8,21 +8,45 @@ from pathlib import Path -def build_plan(config: dict) -> dict: - """Produce paired jobs with a shared baseline identifier for each seed.""" - if config["schema_version"] != 1: +def validate_config(config: dict) -> None: + """Reject plans that silently corrupt the paired experimental control.""" + if type(config.get("schema_version")) is not int or config["schema_version"] != 1: raise ValueError("Unsupported experiment schema version") - seeds = config["seeds"] - arms = config["arms"] - if not seeds or len(set(seeds)) != len(seeds): - raise ValueError("Seeds must be nonempty and unique") + seeds = config.get("seeds") + if (not isinstance(seeds, list) or not seeds + or any(type(seed) is not int or seed < 0 for seed in seeds) + or len(set(seeds)) != len(seeds)): + raise ValueError("Seeds must be nonempty, unique, nonnegative integers") + for key in ("baseline_training_steps", "training_steps_to_do"): + if type(config.get(key)) is not int or config[key] <= 0: + raise ValueError(f"{key} must be a positive integer") + arms = config.get("arms") + if not isinstance(arms, list) or not arms: + raise ValueError("A nonempty list of intervention arms is required") + for arm in arms: + if not isinstance(arm, dict) or not isinstance(arm.get("name"), str) or not arm["name"].strip(): + raise ValueError("Each arm needs a nonempty name") + if type(arm.get("add_neurons")) is not int or arm["add_neurons"] < 0: + raise ValueError("add_neurons must be a nonnegative integer") + if arm.get("sampling") not in ("original", "group_balanced"): + raise ValueError("Unsupported sampling policy") names = [arm["name"] for arm in arms] if len(set(names)) != len(names) or "continue" not in names: raise ValueError("Unique arms including the continued-training control are required") - if config["training_steps_to_do"] <= 0: - raise ValueError("Each arm needs a positive, matched training budget") + control = next(arm for arm in arms if arm["name"] == "continue") + if control["add_neurons"] != 0 or control["sampling"] != "original": + raise ValueError("The continued-training control must not edit capacity or sampling") + if config.get("optimizer_at_fork") != "fresh_identical_optimizer_in_all_arms": + raise ValueError("This proposal requires an identical optimizer-reset policy in all arms") + + +def build_plan(config: dict) -> dict: + """Produce paired jobs with a shared baseline identifier for each seed.""" + validate_config(config) + seeds = config["seeds"] + arms = config["arms"] fingerprint = hashlib.sha256( - json.dumps(config, sort_keys=True).encode() + json.dumps(config, sort_keys=True, allow_nan=False).encode() ).hexdigest() jobs = [] for seed in seeds: diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_plan.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_plan.py new file mode 100644 index 00000000..2a90005a --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_plan.py @@ -0,0 +1,88 @@ +"""Standard-library checks for the planning scaffold; no training or downloads.""" + +import copy +import json +import unittest +from pathlib import Path + +from plan import build_plan + + +class PlanTests(unittest.TestCase): + def setUp(self): + self.config = json.loads(Path(__file__).with_name("experiment.json").read_text()) + + def test_matrix_is_paired_and_not_executed(self): + plan = build_plan(self.config) + self.assertEqual(len(plan["jobs"]), 12) + self.assertEqual(len({job["run_id"] for job in plan["jobs"]}), 12) + for seed in self.config["seeds"]: + jobs = [job for job in plan["jobs"] if job["seed"] == seed] + self.assertEqual(len(jobs), 4) + self.assertEqual(len({job["parent_checkpoint_id"] for job in jobs}), 1) + self.assertEqual({job["training_steps_to_do"] for job in jobs}, {250}) + for job in jobs: + self.assertEqual(job["status"], "planned_not_executed") + self.assertIsNone(job["parent_checkpoint_sha256"]) + self.assertIsNone(job["case_manifest_sha256"]) + + def test_fingerprint_is_stable_and_tracks_config_changes(self): + original = copy.deepcopy(self.config) + first = build_plan(self.config) + self.assertEqual(first, build_plan(dict(reversed(list(self.config.items()))))) + self.assertEqual(self.config, original) + self.config["training_steps_to_do"] += 1 + changed = build_plan(self.config) + self.assertNotEqual(first["config_sha256"], changed["config_sha256"]) + self.assertNotEqual(first["jobs"][0]["run_id"], changed["jobs"][0]["run_id"]) + + def test_invalid_seeds(self): + for seeds in ([], [17, 17], [True], [-1], [1.5], "17"): + with self.subTest(seeds=seeds), self.assertRaises(ValueError): + build_plan({**self.config, "seeds": seeds}) + + def test_invalid_budgets(self): + for key in ("baseline_training_steps", "training_steps_to_do"): + for value in (0, -1, 1.5, True, None): + with self.subTest(key=key, value=value), self.assertRaises(ValueError): + build_plan({**self.config, key: value}) + + def test_invalid_schema(self): + for value in (0, 2, True, "1"): + with self.subTest(value=value), self.assertRaises(ValueError): + build_plan({**self.config, "schema_version": value}) + + def test_missing_or_duplicate_control(self): + arms = self.config["arms"] + for invalid in ([], arms[1:], [*arms, arms[0]], None): + with self.subTest(arms=invalid), self.assertRaises(ValueError): + build_plan({**self.config, "arms": invalid}) + + def test_invalid_arm_fields(self): + for field, value in (("name", ""), ("name", 12), ("add_neurons", -1), + ("add_neurons", True), ("add_neurons", 1.5), + ("sampling", "test_set_balanced")): + config = copy.deepcopy(self.config) + config["arms"][1][field] = value + with self.subTest(field=field, value=value), self.assertRaises(ValueError): + build_plan(config) + + def test_control_cannot_be_an_intervention(self): + for field, value in (("add_neurons", 16), ("sampling", "group_balanced")): + config = copy.deepcopy(self.config) + config["arms"][0][field] = value + with self.subTest(field=field), self.assertRaises(ValueError): + build_plan(config) + + def test_optimizer_policy_must_match(self): + with self.assertRaises(ValueError): + build_plan({**self.config, "optimizer_at_fork": "reset_only_edited_arm"}) + + def test_nonfinite_config_is_rejected(self): + self.config["optimizer"]["lr"] = float("nan") + with self.assertRaises(ValueError): + build_plan(self.config) + + +if __name__ == "__main__": + unittest.main()