diff --git a/docs/index.rst b/docs/index.rst index 8eb61c16..447f379c 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -130,6 +130,7 @@ Weightslab is a Python SDK to inspect, monitor, and edit training behavior for c resource_monitoring experiment_reports + prediction_comparison checkpointing agent export diff --git a/docs/prediction_comparison.rst b/docs/prediction_comparison.rst new file mode 100644 index 00000000..5802a39b --- /dev/null +++ b/docs/prediction_comparison.rst @@ -0,0 +1,63 @@ +Paired prediction comparisons +============================= + +``wl.compare_predictions(control, intervention)`` compares two recorded model +branches on the same held-out cases. It is stateless, torch-free, and does not +require a live ledger, training process, or Studio server. + +Record the baseline checkpoint once, fork the branches, and evaluate both after +equal training budgets. Supply JSON-compatible dictionaries: + +.. code-block:: python + + import weightslab as wl + + # Use real SHA-256 digests from your artifacts, not placeholder values. + provenance = { + "parent_checkpoint_sha256": parent_checkpoint_digest, + "case_manifest_sha256": held_out_manifest_digest, + "preprocessing_sha256": preprocessing_digest, + "evaluation_split": "test", + "training_steps": 250, + "optimizer_policy": "fresh_identical_optimizer_in_all_arms", + } + control = {**provenance, "run_id": "seed17-continue", "cases": control_cases} + intervention = {**provenance, "run_id": "seed17-widen", "cases": edited_cases} + comparison = wl.compare_predictions(control, intervention) + print(comparison["groups"], comparison["corrected_ids"], comparison["regressed_ids"]) + +Each case contains ``sample_id`` (unique nonempty string), ``group`` (nonempty +string), and ``label`` / ``prediction`` (nonnegative integers). Optional +``true_label_margin`` is a finite number or ``None``. Other fields are ignored. +IDs may arrive in different orders; the comparison joins them by identity. + +Safety checks +------------- + +- Both inputs need distinct run IDs and matching checkpoint, case-manifest, + preprocessing, split, training-step and optimizer-policy metadata. +- Digests must be lowercase 64-character SHA-256 strings. Budgets are + nonnegative integer step counts, not elapsed time or epochs. +- Missing or duplicated cases, changed labels/groups, and nonfinite margins + fail. The function never silently intersects two different cohorts. +- ``TypeError`` identifies non-mapping snapshots/cases; ``ValueError`` identifies + other invalid or unpaired input. Inputs are not modified. + +Result and boundaries +--------------------- + +The versioned result includes overall and per-group empirical accuracy, +accuracy differences, counts, corrected/regressed IDs, and sorted per-case +outcomes. Accuracies/differences are **fractions**, not percentages. Missing +margins remain ``None``. Convert differences to percentage points in the UI. + +This validates the **supplied provenance**, not the actual files, training +implementation, or truth of the recorded metadata. It does not certify causal +mechanisms, statistical significance, fairness, or checkpoint replay. Callers +must separately verify identical parent weights, shape/optimizer correctness, +sampling policies and saved-artifact hashes. Dataset-weighted averages and +multi-seed uncertainty are the experiment's responsibility. + +The ``wl-model-editing/hard-example-diagnostics`` example exercises this API +in a four-arm Waterbirds pilot and exports an offline comparison page. That +page is a recorded-run prototype, not a new live Studio screen or RPC. diff --git a/docs/proposals/hard-example-diagnostics.md b/docs/proposals/hard-example-diagnostics.md new file mode 100644 index 00000000..5e0c3a40 --- /dev/null +++ b/docs/proposals/hard-example-diagnostics.md @@ -0,0 +1,178 @@ +# Make one failure understandable — and test whether we can fix it + +**WeightsLab · project brief · updated 8 October 2026** + +**Status:** the first real-data pilot and offline demo have now run. See the +[measured results and presentation script](waterbirds-demo-results-2026-10-08.md). +The roadmap below retains the original proposal; live Studio integration remains 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 + +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/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..57643074 --- /dev/null +++ b/docs/proposals/hard-example-meeting-2026-10-08.md @@ -0,0 +1,129 @@ +# From model editing to explainable experiments + +WeightsLab · 8 October 2026 · meeting handout + +Update: the first pilot has now executed. Use the +[measured results and demo script](waterbirds-demo-results-2026-10-08.md) for +presentation; this handout records the original proposal and decision checklist. + +**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/docs/proposals/waterbirds-demo-results-2026-10-08.md b/docs/proposals/waterbirds-demo-results-2026-10-08.md new file mode 100644 index 00000000..752d97d5 --- /dev/null +++ b/docs/proposals/waterbirds-demo-results-2026-10-08.md @@ -0,0 +1,111 @@ +# Waterbirds: a controlled editing experiment + +8 October 2026 · measured pilot · branch `srini/hard-example-diagnostics` + +## Presentation takeaway + +> We can inspect a recurring failure, try a controlled intervention, and show +> both improvements and regressions. In this experiment, more exposure to the +> rare group helped; adding neurons alone did not. Inspection also surfaced +> known label problems among the hardest examples. + +This is a useful diagnostic result, not evidence that adding capacity reliably +improves a transformer. The backbone was frozen; only the MLP head was edited. + +## Measured results + +Mean test accuracies across paired seeds 17, 29 and 43: + +| Arm | Waterbird on land | Change vs control | Common cases | Overall empirical | +|---|---:|---:|---:|---:| +| Continue | 59.45% | — | 98.19% | 84.32% | +| Add 16 head neurons | 57.11% | −2.34 pp | 96.84% | 82.74% | +| Balanced sampling | 82.87% | +23.42 pp | 95.56% | 90.94% | +| Widen + balanced sampling | 74.71% | +15.26 pp | 93.76% | 86.00% | + +Common cases pool landbird/land and waterbird/water by their test counts. +The balanced arm loses **2.63 percentage points** on this measure, exceeding +the proposed one-point regression budget. No arm satisfies the complete pilot +target. Three seeds are preliminary evidence, not a significance claim. + +Target-group paired gains, seed order 17 / 29 / 43: + +- Widen: −3.27 / −0.62 / −3.12 pp. +- Balanced: +22.90 / +23.99 / +23.36 pp. +- Combined: +14.95 / +12.46 / +18.38 pp. + +The report also records the training-frequency-weighted benchmark average; +it is not interchangeable with the empirical test average above. + +## Protocol and verification + +- Official Waterbirds splits: 4,795 training, 1,199 validation, 5,794 test. +- Frozen torchvision ViT-B/16 `IMAGENET1K_V1`, its pinned preprocessing, + 768-dimensional features, editable 64-unit MLP head, widened to 80 units. +- 500 baseline steps per seed; four branches with 250 steps each. Fresh SGD + without momentum in every branch; learning rate 0.01, training batch size 64. +- Target group chosen by lowest mean baseline validation accuracy, before test + prediction inspection. All official test cases and labels retained. +- Matching sampled batches within capacity pairs; immediate-edit and trained + results recorded separately. Retained weights, dependent layer shapes and + optimizer parameter references checked after every structural edit. +- Two complete runs produced identical final per-case predictions and sampled + training histories. Independently replayed all 12 saved final heads on all + 5,794 held-out cases per head; prediction labels matched exactly. +- Frozen features extracted on NVIDIA L40; torch 2.9.1+cu128 and torchvision + 0.24.1+cu128. Checkpoint, metadata, preprocessing, protocol and runner/API + source hashes are retained in the generated evidence bundle. + +Final protocol hash: +`7902607cf99bc35a06b1518d6e23ba09371b5058b83c7d24e6267005a595557e`. +The run used the working source before its publication commit; the report +explicitly records `source_dirty` and exact runner/comparison source hashes. + +## What to show in five minutes + +1. **Results:** start with balanced sampling and the waterbird/land group. + Point out the +23.42-point gain and −2.63-point common-case tradeoff. +2. **Model editing:** select “Add 16 neurons.” Show 64 → 80 hidden units and + the classifier input growing to 80 while the 768-dimensional backbone stays + fixed. Compare the immediate edit effect with the trained outcome. +3. **Cases:** toggle corrected/regressed examples and inspect predictions, + confidence and margins. Images are post-hoc illustrations; metrics use the + full held-out set. Aggregated charts and single-seed images are labelled. +4. **Attribution:** inspect the same validation image through two trained heads. + These are class-conditioned Integrated Gradients, not changed attention. + Approximation warnings remain visible where convergence checks failed. +5. **History:** open the provenance panel and show the common checkpoint, + intervention, optimizer checks and locked protocol. + +## Feature contribution + +`wl.compare_predictions(control, intervention)` is a new stateless public API. +It validates supplied checkpoint/split/preprocessing/budget/optimizer metadata, +joins stable sample IDs, rejects changed cohorts or labels, and returns +per-group accuracy changes plus corrected and regressed IDs. It does not +silently hide missing cases. Its checks do not independently prove that +supplied provenance is truthful; the runner and replay verifier supply that +additional evidence here. + +The branch also provides the dataset adapter, frozen-feature cache, controlled +runner, saved-checkpoint verifier, numerical attribution checks and offline +interactive demo. [API reference](../prediction_comparison.rst) · +[reproduction commands](../../weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md). + +## Important boundaries + +The authors flag Eastern Towhees, Western Meadowlarks and Western Wood Pewees +as incorrectly labelled waterbirds. Some locked hard validation cases are +Eastern Towhees. The UI flags this; the experiment does not relabel or remove +them after looking at results. [Authors' dataset note](https://github.com/kohpangwei/group_DRO#waterbirds). + +Integrated Gradients used 32–128 Gauss-Legendre samples and a zero normalized +input baseline. Six of the 20 baseline/control/intervention maps failed the +recorded absolute-plus-relative completeness tolerance; warnings are retained. +Do not use small visual differences as causal proof. Method reference: +[Integrated Gradients and completeness checks](https://captum.ai/docs/extension/integrated_gradients). + +The page is an offline recorded-run prototype, not a live Studio screen. No new +RPC, browser/server synchronization, structural transformer editing, full +training-state restore, or reliable real-world rare-animal improvement is +claimed. Those remain separate integration/research work. diff --git a/tests/diagnostics/test_prediction_comparison.py b/tests/diagnostics/test_prediction_comparison.py new file mode 100644 index 00000000..6df38d14 --- /dev/null +++ b/tests/diagnostics/test_prediction_comparison.py @@ -0,0 +1,92 @@ +"""The comparison must never silently pair different evaluation populations.""" + +import copy + +import pytest + +import weightslab as wl + + +@pytest.fixture +def pair(): + control = {"run_id": "control", "parent_checkpoint_sha256": "a" * 64, + "case_manifest_sha256": "b" * 64, "preprocessing_sha256": "c" * 64, + "evaluation_split": "test", "training_steps": 250, "optimizer_policy": "fresh_sgd", + "cases": [{"sample_id": "a", "group": "rare", "label": 1, "prediction": 0, + "true_label_margin": -0.2}, + {"sample_id": "b", "group": "common", "label": 0, "prediction": 0}]} + intervention = copy.deepcopy(control) + intervention["run_id"] = "widen" + intervention["cases"][0].update(prediction=1, true_label_margin=0.3) + intervention["cases"][1]["prediction"] = 1 + return control, intervention + + +def test_corrected_regressed_and_order_independence(pair): + original = copy.deepcopy(pair) + result = wl.compare_predictions(*pair) + assert result["corrected_ids"] == ["a"] + assert result["regressed_ids"] == ["b"] + assert result["overall"]["accuracy_delta"] == 0 + assert result["groups"]["rare"]["accuracy_delta"] == 1 + assert result["groups"]["common"]["accuracy_delta"] == -1 + assert result["cases"][0]["margin_delta"] == pytest.approx(0.5) + assert result["cases"][1]["margin_delta"] is None + assert pair == original + pair[1]["cases"].reverse() + assert wl.compare_predictions(*pair) == result + + +@pytest.mark.parametrize("key,value", [("parent_checkpoint_sha256", "d" * 64), + ("case_manifest_sha256", "d" * 64), + ("preprocessing_sha256", "d" * 64), + ("evaluation_split", "validation"), + ("training_steps", 251), ("optimizer_policy", "keep_state")]) +def test_unpaired_provenance(pair, key, value): + pair[1][key] = value + with pytest.raises(ValueError, match="Unpaired"): + wl.compare_predictions(*pair) + + +@pytest.mark.parametrize("key", ["parent_checkpoint_sha256", "case_manifest_sha256", "run_id", + "preprocessing_sha256", "evaluation_split", "training_steps", "optimizer_policy"]) +def test_missing_provenance(pair, key): + del pair[1][key] + with pytest.raises(ValueError): + wl.compare_predictions(*pair) + + +@pytest.mark.parametrize("key,value", [("label", 0), ("group", "different"), ("sample_id", "new")]) +def test_changed_cohort(pair, key, value): + pair[1]["cases"][0][key] = value + with pytest.raises(ValueError): + wl.compare_predictions(*pair) + + +@pytest.mark.parametrize("key,value", [("prediction", True), ("label", -1), ("sample_id", ""), + ("group", None), ("true_label_margin", float("nan")), + ("true_label_margin", float("inf"))]) +def test_invalid_case(pair, key, value): + pair[1]["cases"][0][key] = value + with pytest.raises(ValueError): + wl.compare_predictions(*pair) + + +def test_empty_missing_and_duplicate_cases(pair): + for cases in ([], pair[1]["cases"][:1], [pair[1]["cases"][0]] * 2): + changed = copy.deepcopy(pair[1]) + changed["cases"] = cases + with pytest.raises(ValueError): + wl.compare_predictions(pair[0], changed) + + +def test_same_run_rejected(pair): + with pytest.raises(ValueError, match="distinct"): + wl.compare_predictions(pair[0], pair[0]) + + +def test_margin_overflow_rejected(pair): + pair[0]["cases"][0]["true_label_margin"] = -1e308 + pair[1]["cases"][0]["true_label_margin"] = 1e308 + with pytest.raises(ValueError, match="overflowed"): + wl.compare_predictions(*pair) diff --git a/weightslab/__init__.py b/weightslab/__init__.py index 759de5ea..2e6a5e44 100644 --- a/weightslab/__init__.py +++ b/weightslab/__init__.py @@ -32,6 +32,7 @@ "seed_everything": (".utils.tools", "seed_everything"), "guard_training_context": (".components.global_monitoring", "guard_training_context"), "guard_testing_context": (".components.global_monitoring", "guard_testing_context"), + "compare_predictions": (".diagnostics", "compare_predictions"), } # Everything re-exported straight from .src (attribute name == export name). for _name in ( @@ -210,6 +211,7 @@ def _clean(v: str) -> str: __credits__ = 'GrayBx' __license__ = 'BSD 2-clause' __all__ = [ + "compare_predictions", "watch_or_edit", "serve", "keep_serving", diff --git a/weightslab/diagnostics.py b/weightslab/diagnostics.py new file mode 100644 index 00000000..7aacf1d8 --- /dev/null +++ b/weightslab/diagnostics.py @@ -0,0 +1,106 @@ +"""Paired, provenance-checked prediction comparisons without a running ledger.""" + +from __future__ import annotations + +import math +import re +from collections.abc import Mapping + + +def _cases(snapshot: Mapping) -> dict: + cases = snapshot.get("cases") + if not isinstance(cases, list) or not cases: + raise ValueError("cases must be a nonempty list") + indexed = {} + for case in cases: + if not isinstance(case, Mapping): + raise TypeError("Each case must be a mapping") + sample_id, group = case.get("sample_id"), case.get("group") + if not isinstance(sample_id, str) or not sample_id or sample_id in indexed: + raise ValueError("sample_id must be a unique, nonempty string") + if not isinstance(group, str) or not group: + raise ValueError("group must be a nonempty string") + for key in ("label", "prediction"): + if type(case.get(key)) is not int or case[key] < 0: + raise ValueError(f"{key} must be a nonnegative integer") + margin = case.get("true_label_margin") + if margin is not None and (type(margin) not in (int, float) or not math.isfinite(margin)): + raise ValueError("true_label_margin must be finite or null") + indexed[sample_id] = case + return indexed + + +def compare_predictions(control: Mapping, intervention: Mapping) -> dict: + """Compare the same held-out cases after equally budgeted branch training. + + Each input is a JSON-compatible snapshot with ``run_id``, ``cases`` and + provenance: ``parent_checkpoint_sha256``, ``case_manifest_sha256``, + ``preprocessing_sha256``, ``evaluation_split``, ``training_steps`` and + ``optimizer_policy``. Each case has a unique string ``sample_id``, string + ``group``, nonnegative integer ``label`` and ``prediction``; optional + ``true_label_margin`` is finite or null. Case order need not match. + + Rejects changed cohorts/labels/groups, missing provenance and mismatched + parents, preprocessing, evaluation splits, budgets or optimizer policies. + This checks supplied provenance, not the checkpoint files or how a model + was trained. The caller remains responsible for recording it truthfully. + + Returns per-group empirical accuracy/delta, corrected/regressed IDs and + per-case outcomes. Accuracies use fractions, not percentages. No statistical + significance, causal mechanism or dataset-weighted average is inferred. + This pure function neither mutates inputs nor reads a global ledger. + """ + if not isinstance(control, Mapping) or not isinstance(intervention, Mapping): + raise TypeError("Both snapshots must be mappings") + hash_keys = ("parent_checkpoint_sha256", "case_manifest_sha256", "preprocessing_sha256") + for snapshot in (control, intervention): + for key in hash_keys: + if not isinstance(snapshot.get(key), str) or not re.fullmatch(r"[0-9a-f]{64}", snapshot[key]): + raise ValueError(f"{key} must be a lowercase SHA-256 digest") + for key in ("run_id", "evaluation_split", "optimizer_policy"): + if not isinstance(snapshot.get(key), str) or not snapshot[key].strip(): + raise ValueError(f"{key} must be a nonempty string") + if type(snapshot.get("training_steps")) is not int or snapshot["training_steps"] < 0: + raise ValueError("training_steps must be a nonnegative integer") + if control["run_id"] == intervention["run_id"]: + raise ValueError("Control and intervention must have distinct run IDs") + for key in (*hash_keys, "evaluation_split", "training_steps", "optimizer_policy"): + if control[key] != intervention[key]: + raise ValueError(f"Unpaired comparison: {key} differs") + before, after = _cases(control), _cases(intervention) + if before.keys() != after.keys(): + raise ValueError("Case cohorts differ; silently intersecting is not allowed") + rows = [] + for sample_id in sorted(before): + left, right = before[sample_id], after[sample_id] + if (left["label"], left["group"]) != (right["label"], right["group"]): + raise ValueError(f"Case label/group changed: {sample_id}") + old, new = left["prediction"] == left["label"], right["prediction"] == right["label"] + outcome = "corrected" if new and not old else "regressed" if old and not new else "unchanged" + left_margin, right_margin = left.get("true_label_margin"), right.get("true_label_margin") + margin_delta = None if left_margin is None or right_margin is None else right_margin - left_margin + if margin_delta is not None and not math.isfinite(margin_delta): + raise ValueError("Margin difference overflowed") + rows.append({"sample_id": sample_id, "group": left["group"], "label": left["label"], + "control_prediction": left["prediction"], "intervention_prediction": right["prediction"], + "control_correct": old, "intervention_correct": new, "outcome": outcome, + "margin_delta": margin_delta}) + + def summarize(subset): + count = len(subset) + old = sum(row["control_correct"] for row in subset) + new = sum(row["intervention_correct"] for row in subset) + return {"count": count, "control_accuracy": old / count, "intervention_accuracy": new / count, + "accuracy_delta": (new - old) / count, + "corrected_count": sum(row["outcome"] == "corrected" for row in subset), + "regressed_count": sum(row["outcome"] == "regressed" for row in subset)} + + return {"schema_version": 1, "artifact_type": "paired_prediction_comparison", + "control_run_id": control["run_id"], "intervention_run_id": intervention["run_id"], + "provenance": {key: control[key] for key in (*hash_keys, "evaluation_split", "training_steps", "optimizer_policy")}, + "overall": summarize(rows), + "groups": {group: summarize([row for row in rows if row["group"] == group]) + for group in sorted({row["group"] for row in rows})}, + "corrected_ids": [row["sample_id"] for row in rows if row["outcome"] == "corrected"], + "regressed_ids": [row["sample_id"] for row in rows if row["outcome"] == "regressed"], + "cases": rows} 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..01ef78dc --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/README.md @@ -0,0 +1,122 @@ +# Hard-example diagnostics + +Controlled experiment and diagnostic prototype 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. 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, + 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. +- `prepare_waterbirds.py`: official metadata adapter and pinned frozen ViT feature cache. +- `run_waterbirds.py`: four arms across three seeds, validation-selected target, + locked test manifest, checkpoint/shape/optimizer checks, per-case predictions + and sampled layer diagnostics. Uses public `wl.compare_predictions`. +- `attribute_waterbirds.py`: full-image Integrated Gradients with checkpoint + prediction parity, shared display scales and numerical-completeness warnings. +- `build_demo.py` + `demo_template.html`: self-contained offline results page. +- `test_experiment.py`: small CPU sampler and attribution numerical tests. + +Synthetic fixture metrics test the plumbing. They are **not Waterbirds, ViT, +or evidence that editing improves real rare cases**. Real results come from +the Waterbirds runner below, not the synthetic fixture. + +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 + +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. + +## Run the real-data experiment + +Requires the repository's torch/torchvision/Pillow/numpy dependencies. A GPU is +recommended for the frozen feature cache and image attribution. The measured +pilot used torch 2.9.1+cu128, torchvision 0.24.1+cu128 and an NVIDIA L40. + +Download and unpack the official `waterbird_complete95_forest2water2` archive +from the [dataset authors](https://github.com/kohpangwei/group_DRO#waterbirds). +Preserve its `metadata.csv`, image paths and official splits. Check the source +dataset terms before redistributing images. Dataset-root below is the directory +containing `metadata.csv`, not its parent. Use fresh cache/run/output paths. + +From the repository root: + +```bash +export WL_NO_TELEMETRY=1 WEIGHTSLAB_OPENCODE_AUTOINSTALL=0 +export CUBLAS_WORKSPACE_CONFIG=:4096:8 PYTHONPATH=. +experiment_dir=weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics + +python "$experiment_dir/prepare_waterbirds.py" \ + --dataset-root /path/to/waterbird_complete95_forest2water2 \ + --output-dir /path/to/new-feature-cache +python "$experiment_dir/run_waterbirds.py" \ + --cache /path/to/new-feature-cache/features.pt --output-dir /path/to/new-run +python "$experiment_dir/attribute_waterbirds.py" \ + --run-dir /path/to/new-run --dataset-root /path/to/waterbird_complete95_forest2water2 + +curl -fL https://cdn.jsdelivr.net/npm/chart.js@4.5.1/dist/chart.umd.min.js \ + --output /path/to/chart-4.5.1.js +python "$experiment_dir/build_demo.py" --run-dir /path/to/new-run \ + --dataset-root /path/to/waterbird_complete95_forest2water2 \ + --chart-js /path/to/chart-4.5.1.js --output /path/to/demo/index.html +python -m http.server 8877 --bind 127.0.0.1 --directory /path/to/demo +``` + +Open `http://localhost:8877/` or open the generated HTML directly. All charts, +sample images and attribution overlays are embedded; the page needs no network +connection. Chart.js is pinned and SHA-256 verified by the builder. The page is +read-only: changing a filter replays saved evidence, not new training or edits. + +The report retains all held-out predictions; preview images are explicitly +post-hoc illustrations. Attribution uses a fixed validation subset, not those +test illustrations. Known dataset label issues and failed numerical attribution +checks remain visible. Branch-local age starts at zero; the parent was trained +for the configured baseline budget. This is learned-head-parameter replay, +not a full training-state/ledger restore. Only load trusted generated caches. + +Run targeted CPU checks: + +```bash +python "$experiment_dir/test_plan.py" +python "$experiment_dir/test_experiment.py" +python -m pytest tests/diagnostics/test_prediction_comparison.py -q +``` + +## Remaining integration work + +- [ ] Register per-case exports with the Studio data ledger; the offline runner + currently exports them to JSON and uses `wl.track_model_signals` for model signals. +- [ ] Add live Studio transport, bounded diagnostic requests and stale-revision handling. +- [ ] Verify pause/edit/resume and browser/server synchronization end to end. +- [ ] Add full-state checkpoint replay beyond learned-head-parameter forks. +- [ ] 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/attribute_waterbirds.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/attribute_waterbirds.py new file mode 100644 index 00000000..505d3e8e --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/attribute_waterbirds.py @@ -0,0 +1,132 @@ +"""Measured class-conditioned input attribution for locked validation cases.""" + +from __future__ import annotations + +import argparse +import base64 +import io +import json +from pathlib import Path + +import numpy as np +import torch +from PIL import Image +from prepare_waterbirds import sha256 +from torch import nn +from torchvision.models import ViT_B_16_Weights, vit_b_16 + +import weightslab as wl + + +def integrated_gradients(model, image, target, steps, batch_size): + """Gauss-Legendre IG from zero normalized input; return signed completeness error.""" + points, weights = np.polynomial.legendre.leggauss(steps) + points = torch.tensor((points + 1) / 2, device=image.device, dtype=image.dtype) + weights = torch.tensor(weights / 2, device=image.device, dtype=image.dtype) + accumulated = torch.zeros_like(image) + for start in range(0, steps, batch_size): + scaled = (points[start:start + batch_size, None, None, None] * image).detach().requires_grad_(True) + output = model(scaled)[:, target] + gradient = torch.autograd.grad(output.sum(), scaled)[0] + accumulated += (gradient * weights[start:start + batch_size, None, None, None]).sum(0, keepdim=True) + attribution = image * accumulated + with torch.no_grad(): + difference = float(model(image)[0, target] - model(torch.zeros_like(image))[0, target]) + delta = float(attribution.sum()) - difference + if not torch.isfinite(attribution).all() or not np.isfinite(delta): + raise AssertionError("Nonfinite attribution") + return attribution.detach(), delta, difference + + +def image_uri(image): + stream = io.BytesIO() + image.save(stream, format="PNG") + return "data:image/png;base64," + base64.b64encode(stream.getvalue()).decode() + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--run-dir", type=Path, required=True) + parser.add_argument("--dataset-root", type=Path, required=True) + parser.add_argument("--config", type=Path, default=Path(__file__).with_name("experiment.json")) + args = parser.parse_args() + report = json.loads((args.run_dir / "report.json").read_text()) + config = json.loads(args.config.read_text()) + if config["backbone"] != report["protocol"]["config"]["backbone"]: + raise ValueError("Attribution backbone must match the recorded experiment") + config.update(root_log_dir=str(args.run_dir / "attribution-ledger"), dataset_root=str(args.dataset_root)) + hp = wl.watch_or_edit(config, flag="hyperparameters", defaults=config) + torch.set_num_threads(int(hp["num_threads"])) + device = str(hp["device"]) + if device == "auto": + device = "cuda" if torch.cuda.is_available() else "cpu" + weights = ViT_B_16_Weights[str(hp["backbone"]["weights"])] + backbone = vit_b_16(weights=weights).eval().to(device) + backbone.heads = nn.Identity() + backbone.requires_grad_(False) + weight_path = Path(torch.hub.get_dir()) / "checkpoints" / weights.url.rsplit("/", 1)[-1] + if sha256(weight_path) != report["feature_provenance"]["weights_sha256"]: + raise ValueError("Attribution backbone weights differ from the feature cache") + seed = report["protocol"]["config"]["seeds"][0] + candidates = report["runs"][f"{seed}-continue"]["demo_validation"] + selected = [] + for group in report["group_names"]: + selected.extend([case for case in candidates if case["group"] == group][:int(hp["attribution_cases_per_group"])]) + names = ["baseline"] + [arm["name"] for arm in report["protocol"]["config"]["arms"]] + result = {"method": "Integrated Gradients / Gauss-Legendre", "seed": seed, + "baseline": "zero in normalized image space (ImageNet mean-color image)", + "target": "official dataset class logit, fixed across all branches; known label issues are not corrected", + "selection": "first locked validation case per group in metadata order; no test cases", + "scale": "shared maximum absolute signed channel-summed attribution across all five snapshots per image", + "protocol_sha256": report["protocol_sha256"], + "settings": {key: value for key, value in config.items() if key.startswith("attribution_")}, "cases": {}} + for case in selected: + with Image.open(Path(str(hp["dataset_root"])) / case["image"]) as image: + tensor = weights.transforms()(image.convert("RGB")).unsqueeze(0).to(device) + crop = tensor[0].detach().cpu() * torch.tensor(weights.transforms().std)[:, None, None] + torch.tensor(weights.transforms().mean)[:, None, None] + crop_array = (crop.permute(1, 2, 0).clamp(0, 1).numpy() * 255).astype("uint8") + attribution_maps, entries = {}, {} + for name in names: + filename = f"baseline-{seed}.pt" if name == "baseline" else f"{seed}-{name}.pt" + state = torch.load(args.run_dir / "checkpoints" / filename, map_location="cpu", weights_only=True) + head = nn.Sequential(nn.Linear(state["0"]["weight"].shape[1], state["0"]["weight"].shape[0]), + nn.LeakyReLU(float(hp["head"]["negative_slope"])), + nn.Linear(state["2"]["weight"].shape[1], state["2"]["weight"].shape[0])) + for path, parameters in state.items(): + head[int(path)].load_state_dict(parameters) + model = nn.Sequential(backbone, head.to(device)).eval().requires_grad_(False) + with torch.no_grad(): + logits = model(tensor)[0].cpu() + expected = next(item for item in (report["baselines"][str(seed)]["validation"]["cases"] if name == "baseline" + else report["runs"][f"{seed}-{name}"]["demo_validation"]) if item["sample_id"] == case["sample_id"]) + if not torch.allclose(logits, torch.tensor(expected["logits"]), atol=1e-4, rtol=1e-4): + raise AssertionError("Full image model disagrees with cached-feature predictions") + steps = int(hp["attribution_steps"]) + while True: + attribution, delta, difference = integrated_gradients(model, tensor, case["label"], steps, + int(hp["attribution_batch_size"])) + tolerance = float(hp["attribution_absolute_tolerance"]) + float(hp["attribution_relative_tolerance"]) * abs(difference) + converged = abs(delta) <= tolerance + if converged or steps >= int(hp["attribution_max_steps"]): + break + steps = min(steps * 2, int(hp["attribution_max_steps"])) + attribution_maps[name] = attribution[0].sum(0).cpu().numpy() + entries[name] = {"steps": steps, "completeness_delta": delta, "logit_difference": difference, + "within_tolerance": converged, "tolerance": tolerance, "prediction": expected["prediction"], + "confidence": expected["confidence"], "true_label_margin": expected["true_label_margin"]} + print(f"Attributed {case['sample_id']} / {name}: steps={steps}, delta={delta:.5f}, passed={converged}", flush=True) + scale = max(float(np.abs(values).max()) for values in attribution_maps.values()) + for name, values in attribution_maps.items(): + strength = np.abs(values) / max(scale, np.finfo(float).eps) + color = np.where((values >= 0)[..., None], np.array([240, 84, 44]), np.array([43, 106, 216])) + overlay = crop_array * (1 - strength[..., None]) + color * strength[..., None] + entries[name]["overlay"] = image_uri(Image.fromarray(overlay.astype("uint8"))) + result["cases"][case["sample_id"]] = {"group": case["group"], "label": case["label"], + "crop": image_uri(Image.fromarray(crop_array)), + "shared_scale": scale, "snapshots": entries} + (args.run_dir / "attribution.json").write_text(json.dumps(result, indent=2, allow_nan=False) + "\n") + wl.clear_all() + + +if __name__ == "__main__": + main() diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/build_demo.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/build_demo.py new file mode 100644 index 00000000..7ed893ad --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/build_demo.py @@ -0,0 +1,74 @@ +"""Build an offline results page from measured report, images and attribution.""" + +from __future__ import annotations + +import argparse +import base64 +import hashlib +import io +import json +from pathlib import Path + +from PIL import Image + +CHART_SHA256 = "48444a82d4edcb5bec0f1965faacdde18d9c17db3063d042abada2f705c9f54a" + + +def build(report, attribution, dataset_root, chart_source, output): + if report["status"] != "completed" or report["artifact_type"] != "waterbirds_controlled_experiment": + raise ValueError("Demo requires a completed real-data report") + if report["protocol_sha256"] != attribution["protocol_sha256"]: + raise ValueError("Attribution and experiment protocols differ") + if hashlib.sha256(chart_source).hexdigest() != CHART_SHA256: + raise ValueError("Expected pinned Chart.js 4.5.1 UMD build") + media_ids = set() + # Deterministic post-hoc illustrations; never used to compute reported metrics. + for comparison in report["comparisons"].values(): + for outcome in ("corrected", "regressed"): + for group in report["group_names"]: + selected = [row["sample_id"] for row in comparison["cases"] if row["outcome"] == outcome and row["group"] == group] + media_ids.update(sorted(selected)[:3]) + media = {} + for sample_id in sorted(media_ids): + path = (dataset_root / sample_id).resolve() + if not path.is_relative_to(dataset_root.resolve()): + raise ValueError("Image path escapes dataset root") + with Image.open(path) as image: + image = image.convert("RGB") + image.thumbnail((224, 224)) + stream = io.BytesIO() + image.save(stream, format="JPEG", quality=80) + media[sample_id] = "data:image/jpeg;base64," + base64.b64encode(stream.getvalue()).decode() + data = {key: report[key] for key in ("protocol", "protocol_sha256", "group_names", "train_group_counts", "split_counts", "created_at")} + data.update(media=media, attribution=attribution, runs={}, comparisons={}) + for run_id, run in report["runs"].items(): + data["runs"][run_id] = {key: run[key] for key in ("seed", "arm", "metrics", "training", "events", "checks", "parent_checkpoint_sha256", "checkpoint_sha256")} + data["runs"][run_id]["preview_cases"] = [case for case in run["cases"] if case["sample_id"] in media] + data["runs"][run_id]["immediate_groups"] = run["immediate_test"]["groups"] + data["runs"][run_id]["baseline_groups"] = run["baseline_test"]["groups"] + for run_id, comparison in report["comparisons"].items(): + data["comparisons"][run_id] = {key: comparison[key] for key in ("overall", "groups")} + data["comparisons"][run_id]["preview_cases"] = [case for case in comparison["cases"] if case["sample_id"] in media] + template = Path(__file__).with_name("demo_template.html").read_text() + serialized = json.dumps(data, allow_nan=False).replace("<", "\\u003c").replace("\u2028", "\\u2028").replace("\u2029", "\\u2029") + html = template.replace("/*__CHART_JS__*/", chart_source.decode()).replace("/*__DATA__*/{}", serialized) + output.parent.mkdir(parents=True, exist_ok=True) + with output.open("x") as stream: + stream.write(html) + print(f"Offline demo ready: {output} ({len(media)} illustrative test images; full split metrics)") + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--run-dir", type=Path, required=True) + parser.add_argument("--dataset-root", type=Path, required=True) + parser.add_argument("--chart-js", type=Path, required=True) + parser.add_argument("--output", type=Path, required=True) + args = parser.parse_args() + build(json.loads((args.run_dir / "report.json").read_text()), + json.loads((args.run_dir / "attribution.json").read_text()), + args.dataset_root, args.chart_js.read_bytes(), args.output) + + +if __name__ == "__main__": + main() diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/demo_template.html b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/demo_template.html new file mode 100644 index 00000000..c8e4300d --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/demo_template.html @@ -0,0 +1,51 @@ + + +WeightsLab · The rare-case experiment + +
+
WEIGHTSLAB / EXPERIMENT 01
MEASURED RESULTS · OFFLINE REPLAY
+
Model editing, with evidence

Does a bigger head solve a rare failure?

A frozen ViT-B/16. One recurring failure. Four controlled interventions.
Inspect what improved, what broke, and what actually changed.

+

Every change updates the comparison.
The control always continues from the same checkpoint.

+
FOCUS-GROUP ACCURACY
COMMON-CASE CHANGE
Pooled landbird/land + waterbird/water
vs matched continued training
CORRECTED / REGRESSED
EVIDENCE SCOPE
held-out images per seed
All groups retained, not a curated metric
+

All groups. All four arms.

Test accuracy (%) · chart always shows the full comparison

Did it learn after the intervention?

Focus-group validation accuracy (%) · not the test set
+
Inspect the edit

The backbone stays frozen. The head is editable.

Shape + optimizer checks passed
ViT-B/1612 frozen transformer blocks
→
768 featuresSame cached CLS representation
→
→

Immediate edit impact

Added parameters are initialized by the current editing API. Retained parameters stay unchanged, but outputs need not. This is not a function-preserving initialization.

Measured hidden-layer weight gradients

Training-step samples after the edit · L2 norm
+
Inspect examples

What got better—and what broke?

ImageOutcome

Illustrations are sampled after evaluation (first IDs among corrections/regressions). They are not a representative sample. Headline metrics use every test image.

+
Look at the same input

Class-conditioned attribution—not attention

Full image → frozen ViT → trained head. Fixed dataset-class target; shared color scale.

Supports target logitOpposes target logitOpacity = attribution magnitude

Actual model crop

Cropped image supplied to the ViT

Continue-training control

Control integrated-gradients overlay

Selected intervention

Intervention integrated-gradients overlay

A heatmap is supporting diagnostic evidence, not proof that a neuron learned a concept. Backbone attention cannot change in this frozen-head experiment. IG baseline: ImageNet mean-color image (zero normalized input).

+
Reproducible by construction

One parent checkpoint. An explainable history.

500 baseline steps→Choose group on validation→Lock protocol→Fork same learned weights→Apply edit / sampling→250 matched steps→Compare every held-out case

Fresh SGD without momentum in every arm. Identical sampled batches within each capacity pair. Every widened branch checks shape propagation, retained parameters and optimizer references before continuing.

Inspect provenance and the feature contribution

wl.compare_predictions(control, intervention) validates supplied checkpoint, split, preprocessing, budget and optimizer provenance, then joins stable case IDs and returns per-group changes plus corrected/regressed examples. It rejects mismatched cohorts rather than silently intersecting them.

Limitations and interpretation

This is a short-budget pilot on a composited-background benchmark, not evidence about natural rare animals. Added capacity did not win here. Balanced sampling improves the selected failure mode but regresses common cases; it does not meet the proposed ≤1-point regression bound. Three seeds are preliminary evidence. No hyperparameter search, significance claim, structural transformer editing or live Studio synchronization is represented by this page.

Charts report empirical group accuracies. The full JSON also records the training-frequency-weighted benchmark average. Target selection used validation, not test. Images here are for research demonstration; refer to dataset terms before redistribution.

+ +
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..3fac8f48 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/experiment.json @@ -0,0 +1,52 @@ +{ + "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, "negative_slope": 0.1}, + "feature_extraction": {"batch_size": 128, "num_workers": 4}, + "num_threads": 4, + "model_signals_every_n_steps": 50, + "demo_cases_per_group": 4, + "attribution_steps": 32, + "attribution_max_steps": 128, + "attribution_batch_size": 8, + "attribution_cases_per_group": 1, + "attribution_relative_tolerance": 0.05, + "attribution_absolute_tolerance": 0.01, + "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..8c23f3d6 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/plan.py @@ -0,0 +1,90 @@ +"""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 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.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") + 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, allow_nan=False).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/prepare_waterbirds.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/prepare_waterbirds.py new file mode 100644 index 00000000..8457fade --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/prepare_waterbirds.py @@ -0,0 +1,106 @@ +"""Cache official Waterbirds images using a pinned, frozen ViT-B/16.""" + +from __future__ import annotations + +import argparse +import csv +import hashlib +import json +from pathlib import Path + +import torch +import torchvision +from PIL import Image +from torch.utils.data import DataLoader, Dataset +from torchvision.models import ViT_B_16_Weights, vit_b_16 + +import weightslab as wl + + +def sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as stream: + for chunk in iter(lambda: stream.read(1024 * 1024), b""): + digest.update(chunk) + return digest.hexdigest() + + +def load_rows(root: Path) -> list[dict]: + rows = [] + with (root / "metadata.csv").open(newline="") as stream: + for row in csv.DictReader(stream): + path = (root / row["img_filename"]).resolve() + if not path.is_relative_to(root.resolve()) or not path.is_file(): + raise ValueError("Image path is missing or escapes the dataset root") + label, place, split = int(row["y"]), int(row["place"]), int(row["split"]) + if label not in (0, 1) or place not in (0, 1) or split not in (0, 1, 2): + raise ValueError("Unexpected Waterbirds label, background or split") + rows.append({"sample_id": row["img_filename"], "image": row["img_filename"], + "label": label, "place": place, "split": split, + "group": f"{label}:{place}"}) + if len({row["sample_id"] for row in rows}) != len(rows): + raise ValueError("Duplicate image identity across official splits") + if {row["split"] for row in rows} != {0, 1, 2}: + raise ValueError("All three official splits are required") + return rows + + +class Images(Dataset): + def __init__(self, root, rows, transform): + self.root, self.rows, self.transform = root, rows, transform + + def __len__(self): + return len(self.rows) + + def __getitem__(self, index): + with Image.open(self.root / self.rows[index]["image"]) as image: + return self.transform(image.convert("RGB")) + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", type=Path, default=Path(__file__).with_name("experiment.json")) + parser.add_argument("--dataset-root", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + args = parser.parse_args() + config = json.loads(args.config.read_text()) + config.update(dataset_root=str(args.dataset_root.resolve()), root_log_dir=str(args.output_dir / "wl")) + args.output_dir.mkdir(parents=True, exist_ok=False) + hp = wl.watch_or_edit(config, flag="hyperparameters", defaults=config) + torch.set_num_threads(int(hp["num_threads"])) + device = str(hp["device"]) + if device == "auto": + device = "cuda" if torch.cuda.is_available() else "cpu" + if str(hp["backbone"]["architecture"]) != "vit_b_16" or not bool(hp["backbone"]["frozen"]): + raise ValueError("This cache supports a frozen vit_b_16 only") + weights = ViT_B_16_Weights[str(hp["backbone"]["weights"])] + model = vit_b_16(weights=weights).eval().to(device) + model.requires_grad_(False) + model.heads = torch.nn.Identity() + rows = load_rows(Path(str(hp["dataset_root"]))) + loader = DataLoader(Images(args.dataset_root, rows, weights.transforms()), shuffle=False, + batch_size=int(hp["feature_extraction"]["batch_size"]), + num_workers=int(hp["feature_extraction"]["num_workers"])) + chunks = [] + with torch.inference_mode(): + for batch in loader: + chunks.append(model(batch.to(device)).cpu()) + print(f"Embedded {sum(len(chunk) for chunk in chunks)}/{len(rows)}", flush=True) + features = torch.cat(chunks) + if features.shape != (len(rows), int(hp["backbone"]["feature_width"])) or not torch.isfinite(features).all(): + raise AssertionError("Feature cache shape or finiteness check failed") + weight_path = Path(torch.hub.get_dir()) / "checkpoints" / weights.url.rsplit("/", 1)[-1] + provenance = {"backbone": str(weights), "preprocessing": str(weights.transforms()), + "weights_sha256": sha256(weight_path), + "metadata_sha256": sha256(args.dataset_root / "metadata.csv"), + "torch_version": torch.__version__, "torchvision_version": torchvision.__version__, + "device": device, "rows": len(rows), "config": config} + provenance["preprocessing_sha256"] = hashlib.sha256(provenance["preprocessing"].encode()).hexdigest() + torch.save({"features": features, "rows": rows, "provenance": provenance}, args.output_dir / "features.pt") + (args.output_dir / "provenance.json").write_text(json.dumps(provenance, indent=2) + "\n") + print(f"Feature cache ready: {args.output_dir / 'features.pt'}", flush=True) + wl.clear_all() + + +if __name__ == "__main__": + main() diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/run_waterbirds.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/run_waterbirds.py new file mode 100644 index 00000000..20070ade --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/run_waterbirds.py @@ -0,0 +1,293 @@ +"""Run the locked four-arm, three-seed Waterbirds head-editing pilot.""" + +from __future__ import annotations + +import argparse +import copy +import hashlib +import json +import math +import subprocess +from datetime import datetime, timezone +from pathlib import Path + +import torch +from plan import validate_config +from prepare_waterbirds import sha256 +from torch import nn + +import weightslab as wl + +GROUP_NAMES = {"0:0": "Landbird / land", "0:1": "Landbird / water", + "1:0": "Waterbird / land", "1:1": "Waterbird / water"} + + +def digest(value): + return hashlib.sha256(json.dumps(value, sort_keys=True, allow_nan=False).encode()).hexdigest() + + +def dump(path, value): + path.write_text(json.dumps(value, indent=2, allow_nan=False) + "\n") + + +def linear_layers(model): + return {item["name"]: item for item in model.get_model_graph()["layers"] if item["type"] == "Linear"} + + +def learned_parameters(model): + return {path: {name: value.detach().cpu().clone() + for name, value in model.get_layer_by_id(info["id"]).named_parameters(recurse=False)} + for path, info in linear_layers(model).items()} + + +def register(config, output, run_id, seed, checkpoint=None): + wl.clear_all() + local = copy.deepcopy(config) + local.update(seed=seed, root_log_dir=str(output / "ledger" / run_id)) + hp = wl.watch_or_edit(local, flag="hyperparameters", defaults=local) + device = str(hp["device"]) + if device == "auto": + device = "cuda" if torch.cuda.is_available() else "cpu" + torch.set_num_threads(int(hp["num_threads"])) + torch.manual_seed(int(hp["seed"])) + raw = nn.Sequential(nn.Linear(int(hp["backbone"]["feature_width"]), int(hp["head"]["hidden_width"])), + nn.LeakyReLU(float(hp["head"]["negative_slope"])), + nn.Linear(int(hp["head"]["hidden_width"]), int(hp["head"]["classes"]))) + raw.task_type = "classification" + if checkpoint is not None: + for path, parameters in checkpoint.items(): + raw[int(path)].load_state_dict(parameters) + model = wl.watch_or_edit(raw, flag="model", device=device, compute_dependencies=True, + forced_model_wrapping=True, skip_previous_auto_load=True, + dummy_input=torch.zeros(1, int(hp["backbone"]["feature_width"]), device=device)) + if str(hp["optimizer"]["name"]) != "SGD" or float(hp["optimizer"]["momentum"]) != 0: + raise ValueError("This pilot intentionally uses fresh SGD without momentum in every arm") + optimizer = wl.watch_or_edit(torch.optim.SGD(model.parameters(), lr=float(hp["optimizer"]["lr"]), + momentum=float(hp["optimizer"]["momentum"])), flag="optimizer") + wl.start_training() + return hp, model, optimizer + + +def evaluate(model, features, rows, batch_size): + model.eval() + device = next(model.parameters()).device + chunks = [] + with wl.guard_testing_context, torch.no_grad(): + for start in range(0, len(rows), batch_size): + chunks.append(model(features[start:start + batch_size].to(device)).detach().cpu()) + logits = torch.cat(chunks) + if len(logits) != len(rows) or not torch.isfinite(logits).all(): + raise AssertionError("Invalid evaluation predictions") + labels = torch.tensor([row["label"] for row in rows]) + losses = nn.functional.cross_entropy(logits, labels, reduction="none") + probabilities = logits.softmax(1) + cases = [{**row, "prediction": int(logits[i].argmax()), "logits": logits[i].tolist(), + "confidence": float(probabilities[i].max()), "loss": float(losses[i]), + "true_label_margin": float(logits[i, row["label"]] - logits[i, 1 - row["label"]])} + for i, row in enumerate(rows)] + groups = {} + for group in GROUP_NAMES: + selected = [case for case in cases if case["group"] == group] + if not selected: + raise ValueError(f"Official evaluation split missing group {group}") + groups[group] = {"count": len(selected), + "accuracy": sum(case["prediction"] == case["label"] for case in selected) / len(selected), + "loss": sum(case["loss"] for case in selected) / len(selected)} + return {"cases": cases, "groups": groups, "model_age": model.get_age(), + "accuracy": sum(case["prediction"] == case["label"] for case in cases) / len(cases)} + + +def schedule(rows, seed, policy, steps, batch_size): + counts = {group: sum(row["group"] == group for row in rows) for group in GROUP_NAMES} + weights = torch.tensor([1 / counts[row["group"]] if policy == "group_balanced" else 1.0 for row in rows], + dtype=torch.float64) + indices = torch.multinomial(weights, steps * batch_size, replacement=True, + generator=torch.Generator().manual_seed(seed)) + return indices.reshape(steps, batch_size) + + +def train(model, optimizer, data, validation, hp, steps, seed, policy): + features, rows = data + device = next(model.parameters()).device + labels = torch.tensor([row["label"] for row in rows], device=device) + features = features.to(device) + batches = schedule(rows, seed, policy, steps, int(hp["data"]["train_loader"]["batch_size"])) + initial_age = model.get_age() + every = int(hp["eval_full_to_train_steps_ratio"]) + tracker = wl.track_model_signals(model, every_n_steps=int(hp["model_signals_every_n_steps"])) + layers = {path: model.get_layer_by_id(info["id"]) for path, info in linear_layers(model).items()} + activations, handles, history = {}, [], [] + for path, layer in layers.items(): + def capture(module, inputs, output, path=path): + activations[path] = {"mean": float(output.detach().mean()), "std": float(output.detach().std(unbiased=False))} + handles.append(layer.register_forward_hook(capture)) + try: + while model.get_age() - initial_age < steps: + offset = model.get_age() - initial_age + record = (offset + 1) % every == 0 or offset + 1 == steps + before = {path: layer.weight.detach().clone() for path, layer in layers.items()} if record else {} + model.train() + indices = batches[offset].to(device) + 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() + gradients = {path: float(layer.weight.grad.norm()) for path, layer in layers.items()} if record else {} + optimizer.step() + complete = True + if not complete or not math.isfinite(float(loss.detach())): + raise AssertionError("Training failed or produced a nonfinite loss") + if record: + diagnostics = {path: {"gradient_norm": gradients[path], "weight_norm": float(layer.weight.detach().norm()), + "update_norm": float((layer.weight.detach() - before[path]).norm()), + "activation": dict(activations[path])} for path, layer in layers.items()} + valid = evaluate(model, *validation, int(hp["data"]["test_loader"]["batch_size"])) + history.append({"step": offset + 1, "model_age": model.get_age(), "training_loss": float(loss.detach()), + "validation_groups": valid["groups"], "layers": diagnostics}) + finally: + tracker.remove() + for handle in handles: + handle.remove() + return {"history": history, "schedule_sha256": hashlib.sha256(batches.numpy().tobytes()).hexdigest(), + "steps_completed": model.get_age() - initial_age} + + +def run(config, cache, output): + validate_config(config) + torch.use_deterministic_algorithms(True) + features, rows = cache["features"], cache["rows"] + if cache["provenance"]["config"]["backbone"] != config["backbone"]: + raise ValueError("Feature cache backbone differs from protocol") + splits = {} + for number, name in enumerate(("train", "validation", "test")): + indices = [i for i, row in enumerate(rows) if row["split"] == number] + splits[name] = features[indices], [rows[i] for i in indices] + output.mkdir(parents=True, exist_ok=False) + checkpoints = output / "checkpoints" + checkpoints.mkdir() + baselines = {} + for seed in config["seeds"]: + hp, model, optimizer = register(config, output, f"baseline-{seed}", seed) + training = train(model, optimizer, splits["train"], splits["validation"], hp, + int(hp["baseline_training_steps"]), seed + 100, "original") + valid = evaluate(model, *splits["validation"], int(hp["data"]["test_loader"]["batch_size"])) + path = checkpoints / f"baseline-{seed}.pt" + torch.save(learned_parameters(model), path) + baselines[str(seed)] = {"validation": valid, "training": training, + "checkpoint_sha256": sha256(path), "graph": model.get_model_graph()} + dump(output / "baseline-validation.json", baselines) + print(f"Baseline {seed}: validation groups {valid['groups']}", flush=True) + target = min(GROUP_NAMES, key=lambda group: sum(baselines[str(seed)]["validation"]["groups"][group]["accuracy"] + for seed in config["seeds"])) + demo_cases = [] + reference = baselines[str(config["seeds"][0])]["validation"]["cases"] + for group in GROUP_NAMES: + ordered = sorted((case for case in reference if case["group"] == group), key=lambda case: (-case["loss"], case["sample_id"])) + demo_cases.extend(case["sample_id"] for case in ordered[:config["demo_cases_per_group"]]) + # Locked BEFORE inspecting any test predictions or choosing an intervention. + protocol = {"created_at": datetime.now(timezone.utc).isoformat(), "config": config, + "target_group": target, "target_selection": "lowest mean baseline validation accuracy across all three seeds", + "common_groups": ["0:0", "1:1"], "demo_validation_ids": demo_cases, + "case_manifest_sha256": digest(splits["test"][1]), + "preprocessing_sha256": cache["provenance"]["preprocessing_sha256"], + "dataset_metadata_sha256": cache["provenance"]["metadata_sha256"], + "threshold_status": "proposed_descriptive_only_not_team_approved"} + dump(output / "locked-protocol.json", protocol) + protocol_hash = sha256(output / "locked-protocol.json") + print(f"Protocol locked: target {GROUP_NAMES[target]}, {protocol_hash}", flush=True) + runs, comparisons = {}, {} + for seed in config["seeds"]: + parent = checkpoints / f"baseline-{seed}.pt" + expected_validation = baselines[str(seed)]["validation"]["cases"] + for arm in config["arms"]: + run_id = f"{seed}-{arm['name']}" + checkpoint = torch.load(parent, map_location="cpu", weights_only=True) + hp, model, optimizer = register(config, output, run_id, seed, checkpoint) + batch_size = int(hp["data"]["test_loader"]["batch_size"]) + if evaluate(model, *splits["validation"], batch_size)["cases"] != expected_validation: + raise AssertionError("Checkpoint prediction parity failed") + before_test = evaluate(model, *splits["test"], batch_size) + graph_before = model.get_model_graph() + layers = linear_layers(model) + events = [] + if arm["add_neurons"]: + model.add_neurons(layers["0"]["id"], count=arm["add_neurons"]) + after = learned_parameters(model) + width = int(hp["head"]["hidden_width"]) + if after["0"]["weight"].shape[0] != width + arm["add_neurons"] or after["2"]["weight"].shape[1] != width + arm["add_neurons"]: + raise AssertionError("Dependency propagation failed") + if not (torch.equal(after["0"]["weight"][:width], checkpoint["0"]["weight"]) + and torch.equal(after["0"]["bias"][:width], checkpoint["0"]["bias"]) + and torch.equal(after["2"]["weight"][:, :width], checkpoint["2"]["weight"]) + and torch.equal(after["2"]["bias"], checkpoint["2"]["bias"])): + raise AssertionError("Widening changed retained parameters") + events.append({"operation": "add_neurons", "layer_path": "0", "live_layer_id": layers["0"]["id"], + "count": arm["add_neurons"], "before_width": width, "after_width": width + arm["add_neurons"], + "reason": "Test added head capacity against matched continued training", "retained_parameters_unchanged": True}) + 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 references stale or missing parameters") + immediate = evaluate(model, *splits["test"], batch_size) + training = train(model, optimizer, splits["train"], splits["validation"], hp, + int(hp["training_steps_to_do"]), seed + 1000, arm["sampling"]) + final = evaluate(model, *splits["test"], batch_size) + valid = evaluate(model, *splits["validation"], batch_size) + final_path = checkpoints / f"{run_id}.pt" + torch.save(learned_parameters(model), final_path) + runs[run_id] = {"run_id": run_id, "seed": seed, "arm": arm, "evaluation_split": "test", + "parent_checkpoint_sha256": sha256(parent), "checkpoint_sha256": sha256(final_path), + "case_manifest_sha256": protocol["case_manifest_sha256"], + "preprocessing_sha256": protocol["preprocessing_sha256"], "protocol_sha256": protocol_hash, + "training_steps": training["steps_completed"], "optimizer_policy": config["optimizer_at_fork"], + "cases": final.pop("cases"), "metrics": final, "baseline_test": before_test, + "immediate_test": immediate, "training": training, + "demo_validation": [case for case in valid["cases"] if case["sample_id"] in demo_cases], + "events": events, "graph_before": graph_before, "graph_after": model.get_model_graph(), + "checks": {"checkpoint_parity": True, "optimizer_binding": True, "finite_training": True}} + if arm["name"] != "continue": + control = runs[f"{seed}-continue"] + comparisons[run_id] = wl.compare_predictions(control, runs[run_id]) + if arm["sampling"] == "original" and control["training"]["schedule_sha256"] != training["schedule_sha256"]: + raise AssertionError("Capacity comparison used mismatched batches") + print(f"Finished {run_id}: test groups {final['groups']}", flush=True) + dump(output / "runs.json", runs) + if runs[f"{seed}-resample"]["training"]["schedule_sha256"] != runs[f"{seed}-widen_resample"]["training"]["schedule_sha256"]: + raise AssertionError("Balanced capacity comparison used mismatched batches") + comparisons[f"{seed}-capacity_balanced"] = wl.compare_predictions(runs[f"{seed}-resample"], runs[f"{seed}-widen_resample"]) + train_counts = {group: sum(row["group"] == group for row in splits["train"][1]) for group in GROUP_NAMES} + for run in runs.values(): + run["metrics"]["training_weighted_accuracy"] = sum(run["metrics"]["groups"][g]["accuracy"] * count for g, count in train_counts.items()) / len(splits["train"][1]) + report = {"schema_version": 1, "artifact_type": "waterbirds_controlled_experiment", "status": "completed", + "created_at": datetime.now(timezone.utc).isoformat(), "protocol": protocol, "protocol_sha256": protocol_hash, + "feature_provenance": cache["provenance"], "group_names": GROUP_NAMES, "train_group_counts": train_counts, + "split_counts": {name: len(data[1]) for name, data in splits.items()}, "baselines": baselines, + "runs": runs, "comparisons": comparisons, + "source_commit": subprocess.check_output(["git", "rev-parse", "HEAD"], text=True).strip(), + "source_dirty": bool(subprocess.check_output(["git", "status", "--porcelain"], text=True).strip()), + "source_sha256": {str(path.relative_to(Path.cwd())): sha256(path) + for path in [Path(__file__).resolve(), Path(wl.__file__).resolve().with_name("diagnostics.py")]}, + "claim_boundary": "Frozen ViT plus editable head; recorded experiment, not live Studio integration or structural attention editing"} + dump(output / "report.json", report) + wl.clear_all() + return report + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--config", type=Path, default=Path(__file__).with_name("experiment.json")) + parser.add_argument("--cache", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, required=True) + args = parser.parse_args() + config = json.loads(args.config.read_text()) + # Trusted local cache produced by prepare_waterbirds.py; never load arbitrary pickles. + cache = torch.load(args.cache, map_location="cpu", weights_only=False) + try: + run(config, cache, args.output_dir) + print(f"Experiment complete: {args.output_dir / 'report.json'}", flush=True) + finally: + wl.clear_all() + + +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}} +} diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_experiment.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_experiment.py new file mode 100644 index 00000000..819a826b --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/test_experiment.py @@ -0,0 +1,45 @@ +"""Small CPU checks for the numerical helpers; no dataset or pretrained download.""" + +import unittest + +import torch +from attribute_waterbirds import integrated_gradients +from run_waterbirds import schedule + + +class ExperimentTests(unittest.TestCase): + def test_original_schedule_is_float_valid_and_reproducible(self): + rows = [{"group": group} for group in ("0:0", "0:1", "1:0", "1:1")] + first = schedule(rows, 17, "original", 10, 8) + self.assertEqual(tuple(first.shape), (10, 8)) + self.assertTrue(torch.equal(first, schedule(rows, 17, "original", 10, 8))) + self.assertFalse(torch.equal(first, schedule(rows, 29, "original", 10, 8))) + self.assertTrue(bool(((first >= 0) & (first < len(rows))).all())) + + def test_balanced_schedule_reweights_unequal_groups(self): + rows = [{"group": group} for group, count in (("0:0", 100), ("0:1", 10), ("1:0", 2), ("1:1", 30)) for _ in range(count)] + sampled = schedule(rows, 17, "group_balanced", 100, 64).flatten().tolist() + for group in ("0:0", "0:1", "1:0", "1:1"): + fraction = sum(rows[i]["group"] == group for i in sampled) / len(sampled) + self.assertLess(abs(fraction - 0.25), 0.03) + + def test_integrated_gradients_matches_linear_ground_truth(self): + model = torch.nn.Sequential(torch.nn.Flatten(), torch.nn.Linear(12, 2)) + model.requires_grad_(False) + image = torch.arange(12, dtype=torch.float32).reshape(1, 3, 2, 2) / 12 + attribution, delta, _ = integrated_gradients(model, image, 1, 8, 3) + expected = image * model[1].weight[1].reshape(1, 3, 2, 2) + self.assertTrue(torch.allclose(attribution, expected, atol=1e-6)) + self.assertLess(abs(delta), 1e-6) + + def test_integrated_gradients_zero_input_has_zero_attribution(self): + model = torch.nn.Sequential(torch.nn.Flatten(), torch.nn.Linear(12, 2), torch.nn.Tanh()) + model.requires_grad_(False) + attribution, delta, difference = integrated_gradients(model, torch.zeros(1, 3, 2, 2), 0, 8, 3) + self.assertEqual(float(attribution.abs().sum()), 0) + self.assertEqual(delta, 0) + self.assertEqual(difference, 0) + + +if __name__ == "__main__": + unittest.main() 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() diff --git a/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/verify_artifacts.py b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/verify_artifacts.py new file mode 100644 index 00000000..0294c73c --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/hard-example-diagnostics/verify_artifacts.py @@ -0,0 +1,68 @@ +"""Replay every saved final head and verify reports against a second full run.""" + +import argparse +import json +from pathlib import Path + +import torch +from prepare_waterbirds import sha256 +from run_waterbirds import digest +from torch import nn + +import weightslab as wl + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--run-dir", type=Path, required=True) + parser.add_argument("--reference-run", type=Path, required=True) + parser.add_argument("--cache", type=Path, required=True) + args = parser.parse_args() + report = json.loads((args.run_dir / "report.json").read_text()) + reference = json.loads((args.reference_run / "report.json").read_text()) + config = report["protocol"]["config"] + hp = wl.watch_or_edit(config, flag="hyperparameters", defaults=config) + torch.set_num_threads(int(hp["num_threads"])) + cache = torch.load(args.cache, map_location="cpu", weights_only=False) + indices = [i for i, row in enumerate(cache["rows"]) if row["split"] == 2] + rows = [cache["rows"][i] for i in indices] + features = cache["features"][indices] + assert digest(rows) == report["protocol"]["case_manifest_sha256"] + assert sha256(args.run_dir / "locked-protocol.json") == report["protocol_sha256"] + checks = {} + for run_id, run in report["runs"].items(): + assert run["cases"] == reference["runs"][run_id]["cases"], f"Predictions changed on repeat: {run_id}" + assert run["metrics"] == reference["runs"][run_id]["metrics"] + assert run["training"] == reference["runs"][run_id]["training"] + path = args.run_dir / "checkpoints" / f"{run_id}.pt" + assert sha256(path) == run["checkpoint_sha256"] + assert sha256(args.run_dir / "checkpoints" / f"baseline-{run['seed']}.pt") == run["parent_checkpoint_sha256"] + state = torch.load(path, map_location="cpu", weights_only=True) + width = int(hp["head"]["hidden_width"]) + run["arm"]["add_neurons"] + head = nn.Sequential(nn.Linear(int(hp["backbone"]["feature_width"]), width), + nn.LeakyReLU(float(hp["head"]["negative_slope"])), + nn.Linear(width, int(hp["head"]["classes"]))) + for layer, parameters in state.items(): + head[int(layer)].load_state_dict(parameters) + with torch.no_grad(): + logits = head(features) + expected = torch.tensor([case["logits"] for case in run["cases"]]) + assert torch.allclose(logits, expected, atol=1e-4, rtol=1e-4), f"Replay failed: {run_id}" + assert logits.argmax(1).tolist() == [case["prediction"] for case in run["cases"]] + for row, case in zip(rows, run["cases"]): + assert all(row[key] == case[key] for key in row) + if run["arm"]["name"] != "continue": + assert wl.compare_predictions(report["runs"][f"{run['seed']}-continue"], run) == report["comparisons"][run_id] + checks[run_id] = {"repeated_predictions_exact": True, "repeated_history_exact": True, + "saved_checkpoint_replayed": True, "max_logit_replay_error": float((logits - expected).abs().max()), + "cases_replayed": len(rows)} + verification = {"status": "passed", "report_sha256": sha256(args.run_dir / "report.json"), + "reference_report_sha256": sha256(args.reference_run / "report.json"), + "protocol_sha256": report["protocol_sha256"], "checks": checks} + (args.run_dir / "verification.json").write_text(json.dumps(verification, indent=2) + "\n") + print(f"Verified {len(checks)} saved heads, {len(rows)} cases each; exact repeated predictions and histories.") + wl.clear_all() + + +if __name__ == "__main__": + main()