diff --git a/docs/model_interaction.rst b/docs/model_interaction.rst index ab1fc609..494e31b7 100644 --- a/docs/model_interaction.rst +++ b/docs/model_interaction.rst @@ -36,9 +36,59 @@ Minimal example log=True, ) +Model graph and editing API +--------------------------- + +Model editing needs dependency information so a change to one layer can be +propagated to connected layers. Enable it when wrapping the model: + +.. code-block:: python + + model = wl.watch_or_edit( + my_model, + flag="model", + dummy_input=example_batch, + compute_dependencies=True, + ) + + graph = model.get_model_graph() + first_layer = model.get_layer_info(graph["layers"][0]["id"]) + +``get_model_graph()`` returns JSON-serializable primitives with a schema +version, model metadata, layers, and directed dependencies. Each dependency is +one of: + +- ``SAME``: both layers expose the same neuron/channel dimension, such as a + convolution followed by batch normalization. +- ``INCOMING``: the source output feeds the destination input, such as one + linear layer feeding another. +- ``REC``: a recursive or skip-connection relationship that must be kept in + sync during structural edits. + +The structural graph omits per-neuron records by default to keep inspection +cheap. Use ``get_model_graph(include_neurons=True)`` or +``get_layer_info(layer_id)`` for learning-rate and frozen-state details. + +The model exposes dependency-aware modifiers: + +.. code-block:: python + + model.freeze_neurons(layer_id=0, neuron_indices=[0, 2]) + model.unfreeze_neurons(layer_id=0, neuron_indices=[0]) + model.reset_neurons(layer_id=0, neuron_indices=[1]) + model.perturb_neurons(layer_id=0, neuron_indices=[2], ratio=0.1) + model.add_neurons(layer_id=0, count=2) + model.prune_neurons(layer_id=0, neuron_indices=[3]) + +Omitting ``neuron_indices`` freezes, unfreezes, resets, or perturbs the whole +layer. Layer ids are stable only for the lifetime of the wrapped model; resolve +them from the graph instead of persisting them as checkpoint identifiers. + Best practices -------------- - Use explicit names for losses/metrics to keep logs readable. - Prefer ``per_sample=True`` for losses when you need hard-example analysis. - Keep model/device arguments explicit to avoid ambiguity in multi-device setups. +- Pause training before structural edits so model and optimizer updates happen + at a safe boundary. diff --git a/tests/model/test_model_graph_api.py b/tests/model/test_model_graph_api.py new file mode 100644 index 00000000..c61b23f8 --- /dev/null +++ b/tests/model/test_model_graph_api.py @@ -0,0 +1,117 @@ +import json +import unittest + +import torch +import torch.nn as nn + +from weightslab.backend import ledgers +from weightslab.backend.model_interface import ModelInterface + + +class TestModelGraphApi(unittest.TestCase): + def tearDown(self): + ledgers.clear_all() + + @staticmethod + def _wrapped_model(): + model = nn.Sequential( + nn.Linear(4, 3), + nn.ReLU(), + nn.Linear(3, 2), + ) + wrapped = ModelInterface( + model, + dummy_input=torch.randn(1, 4), + compute_dependencies=True, + register=False, + skip_previous_auto_load=True, + ) + # Architecture-change hooks update registered optimizers. These focused + # API tests do not register one, so avoid unrelated ledger work. + wrapped._architecture_change_hook_fns = [] + return wrapped + + def test_get_model_graph_returns_serializable_structure(self): + model = self._wrapped_model() + + graph = model.get_model_graph() + + self.assertEqual(graph["schema_version"], 1) + self.assertEqual([layer["name"] for layer in graph["layers"]], ["0", "1", "2"]) + self.assertEqual( + graph["dependencies"], + [ + {"source_layer_id": 0, "target_layer_id": 1, "type": "SAME"}, + {"source_layer_id": 1, "target_layer_id": 2, "type": "INCOMING"}, + ], + ) + self.assertNotIn("neurons", graph["layers"][0]) + json.dumps(graph) + + def test_layer_info_and_modifiers_update_the_live_model(self): + model = self._wrapped_model() + + layer = model.get_layer_info(0) + self.assertEqual(layer["type"], "Linear") + self.assertEqual(layer["input_neurons"], 4) + self.assertEqual(layer["output_neurons"], 3) + self.assertEqual(len(layer["neurons"]), 3) + + model.freeze_neurons(0, [0]) + self.assertTrue(model.get_layer_info(0)["neurons"][0]["frozen"]) + + # Freezing twice is idempotent rather than toggling the state. + model.freeze_neurons(0, [0]) + self.assertTrue(model.get_layer_info(0)["neurons"][0]["frozen"]) + self.assertEqual(model.get_layer_info(0)["operation_counts"]["FREEZE"], 1) + + model.unfreeze_neurons(0, [0]) + self.assertFalse(model.get_layer_info(0)["neurons"][0]["frozen"]) + + model.add_neurons(0, count=2) + self.assertEqual(model.get_layer_info(0)["output_neurons"], 5) + self.assertEqual(model.get_layer_info(2)["input_neurons"], 5) + self.assertEqual(tuple(model(torch.randn(1, 4)).shape), (1, 2)) + + def test_modelling_api_rejects_ambiguous_or_invalid_targets(self): + model = self._wrapped_model() + + with self.assertRaisesRegex(ValueError, "Unknown layer id"): + model.get_layer_info(99) + with self.assertRaisesRegex(ValueError, "cannot be empty"): + model.prune_neurons(0, []) + with self.assertRaisesRegex(ValueError, "outside layer 0"): + model.reset_neurons(0, [3]) + with self.assertRaisesRegex(ValueError, "no learnable weights"): + model.freeze_neurons(1, [0]) + with self.assertRaisesRegex(ValueError, "strictly between 0 and 1"): + model.perturb_neurons(0, [0], ratio=1.0) + + def test_freeze_and_unfreeze_reject_untracked_weight(self): + raw_model = nn.Sequential( + nn.Linear(4, 3), + nn.ReLU(), + nn.Linear(3, 2), + ) + raw_model[0].weight.requires_grad_(False) + model = ModelInterface( + raw_model, + dummy_input=torch.randn(1, 4), + compute_dependencies=True, + register=False, + skip_previous_auto_load=True, + ) + model._architecture_change_hook_fns = [] + + for operation in (model.freeze_neurons, model.unfreeze_neurons): + with self.subTest(operation=operation.__name__): + with self.assertRaisesRegex( + ValueError, "trainable, per-neuron tracked weight" + ): + operation(0, [0]) + + self.assertEqual(model.get_layer_info(0)["operation_counts"]["FREEZE"], 0) + + +if __name__ == "__main__": + unittest.main() diff --git a/weightslab/examples/PyTorch/wl-model-editing/README.md b/weightslab/examples/PyTorch/wl-model-editing/README.md new file mode 100644 index 00000000..f10c328a --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/README.md @@ -0,0 +1,13 @@ +# Model-editing experiments + +This directory contains architecture-specific experiments that validate +WeightsLab model editing against live forward, backward, optimizer, and resumed +training behavior. + +Available experiments: + +- [`vit-model-editing`](vit-model-editing/README.md): ViT-B/16 full-model + compatibility probe and supported editable-head control experiment. + +Add future model families as sibling directories so their compatibility limits +and behavioral assertions remain independently runnable. diff --git a/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/README.md b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/README.md new file mode 100644 index 00000000..3180ec89 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/README.md @@ -0,0 +1,119 @@ +# ViT model-editing experiment + +This experiment tests WeightsLab model editing instead of assuming that an +operation succeeded because the API returned without raising. + +"ViT B-12" is interpreted here as **ViT-Base with 12 transformer encoder +blocks**. The torchvision implementation is `vit_b_16`: Base width/depth with +16x16 input patches. + +## What is tested + +There are two separate scripts because they answer different questions. + +`probe_full_vit.py` is the compatibility gate for wrapping and editing the +complete ViT-B/16 model. It verifies: + +1. Ordinary ViT forward/backward training works. +2. WeightsLab can build its editable dependency graph. +3. A classifier-head freeze can be applied. +4. Training can resume after the edit. + +The current implementation is expected to stop at step 2. Torch FX traces the +model, but WeightsLab dependency mapping reaches `MultiheadAttention`, which +does not expose the neuron-operation interface expected by +`generate_index_maps`. The JSON report records the exact exception. Run without +`--allow-unsupported` when using this as a CI gate. + +`run_head_experiment.py` is the working control experiment. It: + +1. Builds a real 12-block ViT-B/16 backbone. +2. Generates deterministic patterned images without a network download. +3. Caches frozen CLS embeddings from the backbone. +4. Trains a WeightsLab-wrapped MLP classification head. +5. Adds hidden neurons and checks that the downstream input shape changes. +6. Checks that the optimizer is rebuilt with the new parameters. +7. Freezes a neuron and asserts that its gradient is zero. +8. Unfreezes it and asserts that its gradient becomes non-zero. +9. Resumes training and writes all graph, loss, and assertion results to JSON. + +This control is intentionally limited to the head. Editing the 768-dimensional +transformer representation would require synchronized changes across attention +projections, residual paths, positional embeddings, and LayerNorm parameters; +the current dependency engine does not model those relationships reliably. + +## Setup + +From the repository root: + +```bash +bash weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/setup.sh +``` + +The default environment is `.venv-vit-edit`. Override it with +`VIT_EDIT_VENV=/absolute/path` if needed. + +## Run both experiments + +```bash +bash weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_all.sh +``` + +Reports are written under `outputs/vit_model_editing/`, which is gitignored: + +- `full_vit_compatibility_report.json` +- `vit_head_editing_report.json` + +The combined runner allows the known full-ViT incompatibility so the supported +head experiment still runs. For a strict full-model gate: + +```bash +.venv-vit-edit/bin/python \ + weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/probe_full_vit.py +``` + +That command exits non-zero until complete ViT dependency mapping works. + +## Useful variants + +Fast local smoke run (default, no downloads): + +```bash +.venv-vit-edit/bin/python \ + weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_head_experiment.py \ + --image-size 32 --train-samples 16 --eval-samples 8 +``` + +Standard 224x224 input geometry with randomly initialized weights: + +```bash +.venv-vit-edit/bin/python \ + weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_head_experiment.py \ + --image-size 224 +``` + +Pretrained ImageNet backbone (downloads torchvision weights): + +```bash +.venv-vit-edit/bin/python \ + weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_head_experiment.py \ + --image-size 224 --pretrained +``` + +CPU and CUDA are supported. MPS is intentionally excluded because the current +`ModelInterface` device normalization maps non-CUDA devices back to CPU. + +## Reading a result + +A passing head report must have: + +- `status: "passed"` +- `checks.add_propagated: true` +- `checks.optimizer_rebuilt: true` +- `checks.frozen_gradient_norm: 0.0` +- `checks.unfrozen_gradient_norm > 0` +- `checks.forward_after_edit: true` +- `checks.training_resumed: true` + +These are behavioral checks against the live model and optimizer, not only API +shape checks. diff --git a/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/common.py b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/common.py new file mode 100644 index 00000000..04d42923 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/common.py @@ -0,0 +1,140 @@ +"""Shared utilities for the ViT model-editing experiments.""" + +from __future__ import annotations + +import json +import platform +import random +import sys +from pathlib import Path +from typing import Any + +import torch +from torchvision.models import ViT_B_16_Weights, vit_b_16 + + +def seed_everything(seed: int) -> None: + random.seed(seed) + torch.manual_seed(seed) + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) + + +def resolve_device(requested: str) -> torch.device: + if requested == "auto": + return torch.device("cuda" if torch.cuda.is_available() else "cpu") + device = torch.device(requested) + if device.type not in {"cpu", "cuda"}: + raise ValueError( + "This experiment supports CPU and CUDA only because ModelInterface " + "currently normalizes other devices to CPU." + ) + if device.type == "cuda" and not torch.cuda.is_available(): + raise RuntimeError("CUDA was requested but is not available.") + return device + + +def build_vit_b_16( + *, + image_size: int, + num_classes: int, + pretrained: bool, +) -> tuple[torch.nn.Module, int]: + """Build ViT-B/16, the torchvision ViT-Base architecture with 12 blocks.""" + if pretrained: + if image_size != 224: + raise ValueError("Pretrained ViT-B/16 requires --image-size 224.") + model = vit_b_16(weights=ViT_B_16_Weights.DEFAULT) + model.heads.head = torch.nn.Linear(model.hidden_dim, num_classes) + else: + model = vit_b_16( + weights=None, + image_size=image_size, + num_classes=num_classes, + ) + + block_count = len(model.encoder.layers) + if block_count != 12: + raise AssertionError(f"Expected 12 encoder blocks, found {block_count}.") + return model, block_count + + +def make_pattern_images( + *, + samples: int, + image_size: int, + num_classes: int, + seed: int, + normalize: bool, +) -> tuple[torch.Tensor, torch.Tensor]: + """Create deterministic, learnable image patterns without downloading data.""" + if samples < num_classes: + raise ValueError("samples must be at least num_classes.") + + generator = torch.Generator().manual_seed(seed) + labels = torch.arange(samples) % num_classes + images = torch.rand(samples, 3, image_size, image_size, generator=generator) * 0.08 + band = max(1, image_size // 4) + + for index, label_tensor in enumerate(labels): + label = int(label_tensor) + channel = label % 3 + position = (label // 3) % 4 + if label % 2 == 0: + start = min(position * band, image_size - band) + images[index, channel, start : start + band, :] += 0.85 + else: + start = min(position * band, image_size - band) + images[index, channel, :, start : start + band] += 0.85 + + images.clamp_(0.0, 1.0) + if normalize: + mean = torch.tensor([0.485, 0.456, 0.406]).view(1, 3, 1, 1) + std = torch.tensor([0.229, 0.224, 0.225]).view(1, 3, 1, 1) + images = (images - mean) / std + return images, labels.long() + + +@torch.no_grad() +def extract_embeddings( + backbone: torch.nn.Module, + images: torch.Tensor, + *, + batch_size: int, + device: torch.device, +) -> torch.Tensor: + backbone.eval() + outputs = [] + for start in range(0, len(images), batch_size): + batch = images[start : start + batch_size].to(device) + outputs.append(backbone(batch).cpu()) + return torch.cat(outputs) + + +def environment_info(device: torch.device) -> dict[str, Any]: + import torchvision + + import weightslab + + return { + "python": sys.version.split()[0], + "platform": platform.platform(), + "torch": torch.__version__, + "torchvision": torchvision.__version__, + "weightslab": getattr(weightslab, "__version__", "unknown"), + "device": str(device), + } + + +def write_report(path: Path, report: dict[str, Any]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(report, indent=2, sort_keys=True) + "\n") + + +def optimizer_parameter_ids(optimizer: Any) -> set[int]: + raw_optimizer = getattr(optimizer, "optimizer", optimizer) + return { + id(parameter) + for group in raw_optimizer.param_groups + for parameter in group["params"] + } diff --git a/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/probe_full_vit.py b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/probe_full_vit.py new file mode 100755 index 00000000..1a6b1c17 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/probe_full_vit.py @@ -0,0 +1,158 @@ +#!/usr/bin/env python3 +"""Compatibility gate for editing the complete torchvision ViT-B/16 model.""" + +from __future__ import annotations + +import argparse +import traceback +from pathlib import Path + +import torch +from common import ( + build_vit_b_16, + environment_info, + make_pattern_images, + resolve_device, + seed_everything, + write_report, +) +from torch import nn + +import weightslab as wl +from weightslab.backend import ledgers +from weightslab.components.global_monitoring import ( + guard_training_context, + pause_controller, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Probe whether WeightsLab can train and edit a complete ViT-B/16." + ) + parser.add_argument("--output-dir", type=Path, default=Path("outputs/vit_model_editing")) + parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"]) + parser.add_argument("--image-size", type=int, default=32) + parser.add_argument("--num-classes", type=int, default=4) + parser.add_argument("--batch-size", type=int, default=2) + parser.add_argument("--seed", type=int, default=17) + parser.add_argument("--pretrained", action="store_true") + parser.add_argument( + "--allow-unsupported", + action="store_true", + help="Return success while still recording an unsupported result.", + ) + return parser.parse_args() + + +def one_raw_training_step(model, images, labels, optimizer) -> float: + model.train() + optimizer.zero_grad(set_to_none=True) + loss = nn.functional.cross_entropy(model(images), labels) + loss.backward() + optimizer.step() + return float(loss.detach()) + + +def main() -> int: + args = parse_args() + seed_everything(args.seed) + device = resolve_device(args.device) + report_path = args.output_dir / "full_vit_compatibility_report.json" + report = { + "status": "running", + "scope": "Complete torchvision ViT-B/16 model", + "config": vars(args) | {"output_dir": str(args.output_dir)}, + "environment": environment_info(device), + "phases": {}, + } + + try: + images, labels = make_pattern_images( + samples=max(args.batch_size, args.num_classes), + image_size=args.image_size, + num_classes=args.num_classes, + seed=args.seed, + normalize=args.pretrained, + ) + images = images[: args.batch_size].to(device) + labels = labels[: args.batch_size].to(device) + raw_model, block_count = build_vit_b_16( + image_size=args.image_size, + num_classes=args.num_classes, + pretrained=args.pretrained, + ) + raw_model.to(device) + baseline_optimizer = torch.optim.SGD(raw_model.parameters(), lr=1e-3) + baseline_loss = one_raw_training_step( + raw_model, images, labels, baseline_optimizer + ) + report["phases"]["baseline_training"] = { + "status": "passed", + "loss": baseline_loss, + "encoder_blocks": block_count, + } + + ledgers.clear_all() + args.output_dir.mkdir(parents=True, exist_ok=True) + wl.watch_or_edit( + { + "root_log_dir": str(args.output_dir / "full_vit_weightslab_state"), + "skip_checkpoint_load": True, + "experiment_dump_to_train_steps_ratio": 0, + }, + flag="hyperparameters", + ) + model = wl.watch_or_edit( + raw_model, + flag="model", + device=str(device), + dummy_input=images[:1], + compute_dependencies=True, + forced_model_wrapping=True, + skip_previous_auto_load=True, + ) + graph = model.get_model_graph() + report["phases"]["weightslab_graph"] = { + "status": "passed", + "layers": len(graph["layers"]), + "dependencies": len(graph["dependencies"]), + } + + head = next(layer for layer in graph["layers"] if layer["name"] == "heads.head") + optimizer = wl.watch_or_edit( + torch.optim.SGD(model.parameters(), lr=1e-3), + flag="optimizer", + ) + pause_controller.resume(force=True) + model.freeze_neurons(head["id"], [0]) + with guard_training_context: + optimizer.zero_grad(set_to_none=True) + loss = nn.functional.cross_entropy(model(images), labels) + loss.backward() + optimizer.step() + report["phases"]["weightslab_edit_and_resume"] = { + "status": "passed", + "loss": float(loss.detach()), + } + report["status"] = "passed" + print(f"PASS: complete ViT editing is supported; report written to {report_path}") + return 0 + except Exception as exc: # noqa: BLE001 - compatibility probes must report any failure + report.update( + { + "status": "unsupported", + "error": {"type": type(exc).__name__, "message": str(exc)}, + "traceback": traceback.format_exc(), + } + ) + print(f"UNSUPPORTED: {type(exc).__name__}: {exc}") + print(f"Report written to {report_path}") + return 0 if args.allow_unsupported else 1 + finally: + write_report(report_path, report) + ledgers.clear_all() + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_all.sh b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_all.sh new file mode 100755 index 00000000..c165a64b --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_all.sh @@ -0,0 +1,21 @@ +#!/usr/bin/env bash +set -euo pipefail + +script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +repo_root="$(cd "${script_dir}/../../../../.." && pwd)" +venv_path="${VIT_EDIT_VENV:-${repo_root}/.venv-vit-edit}" +python_bin="${venv_path}/bin/python" +output_dir="${VIT_EDIT_OUTPUT:-${repo_root}/outputs/vit_model_editing}" + +export WL_NO_TELEMETRY="${WL_NO_TELEMETRY:-1}" +export WEIGHTSLAB_LOG_LEVEL="${WEIGHTSLAB_LOG_LEVEL:-INFO}" +export WEIGHTSLAB_DISABLE_WATCHDOGS="${WEIGHTSLAB_DISABLE_WATCHDOGS:-1}" + +"${python_bin}" "${script_dir}/probe_full_vit.py" \ + --output-dir "${output_dir}" \ + --allow-unsupported + +"${python_bin}" "${script_dir}/run_head_experiment.py" \ + --output-dir "${output_dir}" + +echo "Experiment reports: ${output_dir}" diff --git a/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_head_experiment.py b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_head_experiment.py new file mode 100755 index 00000000..b8df88ea --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/run_head_experiment.py @@ -0,0 +1,376 @@ +#!/usr/bin/env python3 +"""End-to-end ViT feature extractor + editable MLP head experiment.""" + +from __future__ import annotations + +import argparse +import math +import traceback +from pathlib import Path + +import torch +from common import ( + build_vit_b_16, + environment_info, + extract_embeddings, + make_pattern_images, + optimizer_parameter_ids, + resolve_device, + seed_everything, + write_report, +) +from torch import nn + +import weightslab as wl +from weightslab.backend import ledgers +from weightslab.components.global_monitoring import ( + guard_testing_context, + guard_training_context, + pause_controller, +) + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Train a ViT-B/16 feature pipeline, edit its MLP head, and resume training." + ) + parser.add_argument("--output-dir", type=Path, default=Path("outputs/vit_model_editing")) + parser.add_argument("--device", default="auto", choices=["auto", "cpu", "cuda"]) + parser.add_argument("--image-size", type=int, default=32) + parser.add_argument("--num-classes", type=int, default=4) + parser.add_argument("--train-samples", type=int, default=16) + parser.add_argument("--eval-samples", type=int, default=8) + parser.add_argument("--batch-size", type=int, default=4) + parser.add_argument("--hidden-size", type=int, default=16) + parser.add_argument("--add-neurons", type=int, default=2) + parser.add_argument("--before-steps", type=int, default=4) + parser.add_argument("--after-steps", type=int, default=4) + parser.add_argument("--learning-rate", type=float, default=0.05) + parser.add_argument("--seed", type=int, default=17) + parser.add_argument("--pretrained", action="store_true") + return parser.parse_args() + + +def train_steps( + model, + optimizer, + features: torch.Tensor, + labels: torch.Tensor, + *, + steps: int, + batch_size: int, + capture_layer=None, + capture_neuron: int = 0, +) -> tuple[list[float], float | None]: + model.train() + criterion = nn.CrossEntropyLoss() + losses: list[float] = [] + captured_gradient_norm = None + + for step in range(steps): + start = (step * batch_size) % len(features) + indices = torch.arange(start, start + batch_size) % len(features) + with guard_training_context: + optimizer.zero_grad(set_to_none=True) + logits = model(features[indices]) + loss = criterion(logits, labels[indices]) + loss.backward() + if capture_layer is not None and step == 0: + gradient = capture_layer.weight.grad + if gradient is None: + raise AssertionError("Expected a weight gradient, but found None.") + captured_gradient_norm = float(gradient[capture_neuron].norm().item()) + optimizer.step() + losses.append(float(loss.detach())) + return losses, captured_gradient_norm + + +def active_neuron_for_batch( + model, + optimizer, + features: torch.Tensor, + labels: torch.Tensor, + *, + batch_size: int, + layer, +) -> tuple[int, float]: + """Run one ordinary step and return a neuron proven active on that batch.""" + model.train() + with guard_training_context: + optimizer.zero_grad(set_to_none=True) + logits = model(features[:batch_size]) + loss = nn.functional.cross_entropy(logits, labels[:batch_size]) + loss.backward() + gradient = layer.weight.grad + if gradient is None: + raise AssertionError("Expected hidden-layer gradients, but found None.") + row_norms = gradient.norm(dim=1) + neuron = int(row_norms.argmax().item()) + norm = float(row_norms[neuron].item()) + if norm <= 1e-12: + raise AssertionError("No active hidden neuron was found for the test batch.") + optimizer.step() + return neuron, norm + + +@torch.no_grad() +def evaluate(model, features: torch.Tensor, labels: torch.Tensor) -> dict[str, float]: + model.eval() + with guard_testing_context: + logits = model(features) + loss = nn.functional.cross_entropy(logits, labels) + accuracy = (logits.argmax(dim=1) == labels).float().mean() + return {"loss": float(loss), "accuracy": float(accuracy)} + + +def find_edit_target(graph: dict, hidden_size: int) -> tuple[int, int]: + linear_layers = [layer for layer in graph["layers"] if layer["type"] == "Linear"] + source = next( + (layer for layer in linear_layers if layer["output_neurons"] == hidden_size), + None, + ) + if source is None: + raise AssertionError("Could not identify the editable hidden Linear layer.") + + downstream = next( + ( + layer + for layer in linear_layers + if layer["id"] != source["id"] and layer["input_neurons"] == hidden_size + ), + None, + ) + if downstream is None: + raise AssertionError("Could not identify the downstream classifier Linear layer.") + return source["id"], downstream["id"] + + +def main() -> int: + args = parse_args() + print("[1/7] Resolving environment and synthetic data", flush=True) + seed_everything(args.seed) + device = resolve_device(args.device) + report_path = args.output_dir / "vit_head_editing_report.json" + report = { + "status": "running", + "scope": "ViT-B/16 frozen feature extractor with a WeightsLab-editable MLP head", + "config": vars(args) | {"output_dir": str(args.output_dir)}, + "environment": environment_info(device), + "checks": {}, + } + + try: + train_images, train_labels = make_pattern_images( + samples=args.train_samples, + image_size=args.image_size, + num_classes=args.num_classes, + seed=args.seed, + normalize=args.pretrained, + ) + eval_images, eval_labels = make_pattern_images( + samples=args.eval_samples, + image_size=args.image_size, + num_classes=args.num_classes, + seed=args.seed + 1, + normalize=args.pretrained, + ) + + backbone, block_count = build_vit_b_16( + image_size=args.image_size, + num_classes=args.num_classes, + pretrained=args.pretrained, + ) + embedding_size = backbone.hidden_dim + print(f"[2/7] Extracting embeddings with ViT-B/16 ({block_count} blocks)", flush=True) + backbone.heads = nn.Identity() + backbone.requires_grad_(False) + backbone.to(device) + train_features = extract_embeddings( + backbone, + train_images, + batch_size=args.batch_size, + device=device, + ) + eval_features = extract_embeddings( + backbone, + eval_images, + batch_size=args.batch_size, + device=device, + ) + del backbone + if device.type == "cuda": + torch.cuda.empty_cache() + + print("[3/7] Wrapping the editable MLP head", flush=True) + args.output_dir.mkdir(parents=True, exist_ok=True) + ledgers.clear_all() + wl.watch_or_edit( + { + "root_log_dir": str(args.output_dir / "weightslab_state"), + "skip_checkpoint_load": True, + "experiment_dump_to_train_steps_ratio": 0, + }, + flag="hyperparameters", + ) + raw_head = nn.Sequential( + nn.Linear(embedding_size, args.hidden_size), + nn.ReLU(), + nn.Linear(args.hidden_size, args.num_classes), + ) + model = wl.watch_or_edit( + raw_head, + flag="model", + device="cpu", + dummy_input=train_features[:1], + compute_dependencies=True, + forced_model_wrapping=True, + skip_previous_auto_load=True, + ) + optimizer = wl.watch_or_edit( + torch.optim.SGD(model.parameters(), lr=args.learning_rate), + flag="optimizer", + ) + # This is a finite, headless tensor experiment, so no Studio/serve flow + # exists to resume the training guard for us. + pause_controller.resume(force=True) + + graph_before = model.get_model_graph(include_neurons=True) + print("[4/7] Training the baseline head", flush=True) + source_id, downstream_id = find_edit_target(graph_before, args.hidden_size) + before_metrics = evaluate(model, eval_features, eval_labels) + before_losses, _ = train_steps( + model, + optimizer, + train_features, + train_labels, + steps=args.before_steps, + batch_size=args.batch_size, + ) + + model.add_neurons(source_id, count=args.add_neurons) + print(f"[5/7] Added {args.add_neurons} hidden neurons; validating propagation", flush=True) + graph_after_add = model.get_model_graph(include_neurons=True) + source_after = model.get_layer_info(source_id) + downstream_after = model.get_layer_info(downstream_id) + expected_hidden = args.hidden_size + args.add_neurons + if source_after["output_neurons"] != expected_hidden: + raise AssertionError("The hidden layer did not gain the requested neurons.") + if downstream_after["input_neurons"] != expected_hidden: + raise AssertionError("The downstream classifier input was not propagated.") + + model_parameter_ids = {id(parameter) for parameter in model.parameters()} + optimizer_ids_after_add = optimizer_parameter_ids(optimizer) + if model_parameter_ids != optimizer_ids_after_add: + raise AssertionError("The optimizer does not reference every post-edit parameter.") + + after_add_losses, _ = train_steps( + model, + optimizer, + train_features, + train_labels, + steps=args.after_steps, + batch_size=args.batch_size, + ) + + source_layer = model.get_layer_by_id(source_id) + freeze_neuron, pre_freeze_gradient_norm = active_neuron_for_batch( + model, + optimizer, + train_features, + train_labels, + batch_size=args.batch_size, + layer=source_layer, + ) + model.freeze_neurons(source_id, [freeze_neuron]) + print("[6/7] Validating frozen and unfrozen gradients", flush=True) + if not model.get_layer_info(source_id)["neurons"][freeze_neuron]["frozen"]: + raise AssertionError(f"Neuron {freeze_neuron} was not reported frozen.") + source_layer = model.get_layer_by_id(source_id) + _, frozen_gradient_norm = train_steps( + model, + optimizer, + train_features, + train_labels, + steps=1, + batch_size=args.batch_size, + capture_layer=source_layer, + capture_neuron=freeze_neuron, + ) + if frozen_gradient_norm is None or frozen_gradient_norm > 1e-12: + raise AssertionError( + f"Frozen neuron gradient should be zero, got {frozen_gradient_norm}." + ) + + model.unfreeze_neurons(source_id, [freeze_neuron]) + if model.get_layer_info(source_id)["neurons"][freeze_neuron]["frozen"]: + raise AssertionError(f"Neuron {freeze_neuron} was not reported unfrozen.") + source_layer = model.get_layer_by_id(source_id) + _, unfrozen_gradient_norm = train_steps( + model, + optimizer, + train_features, + train_labels, + steps=1, + batch_size=args.batch_size, + capture_layer=source_layer, + capture_neuron=freeze_neuron, + ) + if unfrozen_gradient_norm is None or unfrozen_gradient_norm <= 1e-12: + raise AssertionError( + f"Unfrozen neuron gradient should be non-zero, got {unfrozen_gradient_norm}." + ) + + final_metrics = evaluate(model, eval_features, eval_labels) + print("[7/7] Writing experiment report", flush=True) + all_losses = before_losses + after_add_losses + if not all(math.isfinite(loss) for loss in all_losses): + raise AssertionError("Training produced a non-finite loss.") + + report.update( + { + "status": "passed", + "architecture": { + "name": "torchvision vit_b_16", + "encoder_blocks": block_count, + "embedding_size": embedding_size, + "editable_scope": "MLP classification head", + }, + "graph_before": graph_before, + "graph_after_add": graph_after_add, + "metrics": { + "before_training": before_metrics, + "after_training_and_edits": final_metrics, + "training_losses": all_losses, + }, + "checks": { + "add_propagated": True, + "optimizer_rebuilt": True, + "tested_neuron": freeze_neuron, + "pre_freeze_gradient_norm": pre_freeze_gradient_norm, + "frozen_gradient_norm": frozen_gradient_norm, + "unfrozen_gradient_norm": unfrozen_gradient_norm, + "forward_after_edit": True, + "training_resumed": True, + }, + } + ) + print(f"PASS: report written to {report_path}") + return 0 + except Exception as exc: # noqa: BLE001 - experiment reports must capture any failure + report.update( + { + "status": "failed", + "error": {"type": type(exc).__name__, "message": str(exc)}, + "traceback": traceback.format_exc(), + } + ) + print(f"FAIL: {type(exc).__name__}: {exc}") + print(f"Report written to {report_path}") + return 1 + finally: + write_report(report_path, report) + ledgers.clear_all() + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/setup.sh b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/setup.sh new file mode 100755 index 00000000..d866a447 --- /dev/null +++ b/weightslab/examples/PyTorch/wl-model-editing/vit-model-editing/setup.sh @@ -0,0 +1,14 @@ +#!/usr/bin/env bash +set -euo pipefail + +script_dir="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" +repo_root="$(cd "${script_dir}/../../../../.." && pwd)" +venv_path="${VIT_EDIT_VENV:-${repo_root}/.venv-vit-edit}" +python_bin="${PYTHON_BIN:-python3}" + +"${python_bin}" -m venv "${venv_path}" +"${venv_path}/bin/python" -m pip install --upgrade pip +"${venv_path}/bin/python" -m pip install --editable "${repo_root}" + +echo "Environment ready: ${venv_path}" +echo "Run: ${venv_path}/bin/python ${script_dir}/run_head_experiment.py" diff --git a/weightslab/models/model_with_ops.py b/weightslab/models/model_with_ops.py index 322abcff..72ab40dc 100755 --- a/weightslab/models/model_with_ops.py +++ b/weightslab/models/model_with_ops.py @@ -3,9 +3,10 @@ from torch import nn from enum import Enum -from typing import List, Set, Optional, Callable, Dict, Any +from typing import Iterable, List, Set, Optional, Callable, Dict, Any from weightslab.components.tracking import TrackingMode +from weightslab.modules.neuron_ops import ArchitectureNeuronsOpType from weightslab.utils.tools import get_children from weightslab.utils.modules_dependencies import _ModulesDependencyManager, DepType @@ -104,7 +105,295 @@ def get_name(self): return self.name def get_layer_by_id(self, layer_id: int): - return self._dep_manager.id_2_layer[layer_id] + layer = self._dep_manager.id_2_layer.get(layer_id) + if layer is None: + available = sorted(self._dep_manager.id_2_layer) + raise ValueError( + f"Unknown layer id {layer_id}. Available layer ids: {available}. " + "Wrap the model with compute_dependencies=True before using " + "the modelling API." + ) + return layer + + def _layer_names_by_id(self) -> Dict[int, str]: + """Return qualified PyTorch names for the registered editable layers.""" + root = getattr(self, "model", None) + if not isinstance(root, nn.Module): + root = self + + names = {} + for qualified_name, module in root.named_modules(): + get_module_id = getattr(module, "get_module_id", None) + if not callable(get_module_id): + continue + try: + names[get_module_id()] = qualified_name or "" + except (AttributeError, TypeError): + continue + return names + + @staticmethod + def _neuron_count(layer: nn.Module, attribute: str) -> Optional[int]: + getter = getattr(layer, "get_neurons", None) + try: + value = getter(attribute) if callable(getter) else getattr(layer, attribute, None) + except (AttributeError, TypeError, ValueError): + value = getattr(layer, attribute, None) + return int(value) if isinstance(value, (int, float)) else None + + def get_layer_info( + self, + layer_id: int, + *, + include_neurons: bool = True, + ) -> Dict[str, Any]: + """Return a JSON-serializable description of one editable layer. + + Layer ids are stable for the lifetime of a wrapped model. They are not + checkpoint identifiers and may change when a new model is wrapped. + """ + return self._build_layer_info( + layer_id, + include_neurons=include_neurons, + layer_names=self._layer_names_by_id(), + ) + + def _build_layer_info( + self, + layer_id: int, + *, + include_neurons: bool, + layer_names: Dict[int, str], + ) -> Dict[str, Any]: + layer = self.get_layer_by_id(layer_id) + input_neurons = self._neuron_count(layer, "in_neurons") + output_neurons = self._neuron_count(layer, "out_neurons") + + lr_overrides = getattr(layer, "neuron_2_lr", {}) + weight_lrs = lr_overrides.get("weight", {}) if lr_overrides else {} + neuron_count = output_neurons or 0 + neuron_lrs = [float(weight_lrs.get(index, 1.0)) for index in range(neuron_count)] + frozen_neurons = {index for index, lr in enumerate(neuron_lrs) if lr == 0.0} + + layer_info = { + "id": int(layer_id), + "name": layer_names.get(layer_id, f"layer_{layer_id}"), + "type": getattr(layer, "module_name", layer.__class__.__name__), + "input_neurons": input_neurons, + "output_neurons": output_neurons, + "parameter_count": sum( + parameter.numel() for parameter in layer.parameters(recurse=False) + ), + "frozen_neuron_count": len(frozen_neurons), + "operation_counts": dict(getattr(layer, "operation_age", {})), + } + if include_neurons: + layer_info["neurons"] = [ + { + "id": index, + "learning_rate": neuron_lrs[index], + "frozen": index in frozen_neurons, + } + for index in range(neuron_count) + ] + return layer_info + + def get_model_graph(self, *, include_neurons: bool = False) -> Dict[str, Any]: + """Return the editable model graph as JSON-serializable primitives. + + The default response is intentionally structural and cheap. Pass + ``include_neurons=True`` when per-neuron learning-rate/frozen state is + needed. + """ + layer_ids = sorted(self._dep_manager.id_2_layer) + if not layer_ids: + raise RuntimeError( + "Model graph information is unavailable. Wrap the model with " + "wl.watch_or_edit(..., flag='model', compute_dependencies=True)." + ) + + dependency_edges = set() + for dependency_type, parents in self._dep_manager.dependency_2_id_2_id.items(): + for source_id, target_ids in parents.items(): + dependency_edges.update( + (int(source_id), int(target_id), dependency_type.value) + for target_id in target_ids + ) + dependencies = [ + { + "source_layer_id": source_id, + "target_layer_id": target_id, + "type": dependency_type, + } + for source_id, target_id, dependency_type in sorted(dependency_edges) + ] + layer_names = self._layer_names_by_id() + + return { + "schema_version": 1, + "name": self.get_name(), + "age": self.get_age(), + "layers": [ + self._build_layer_info( + layer_id, + include_neurons=include_neurons, + layer_names=layer_names, + ) + for layer_id in layer_ids + ], + "dependencies": dependencies, + } + + def _validated_neuron_indices( + self, + layer_id: int, + neuron_indices: Optional[Iterable[int] | int], + *, + allow_empty: bool, + ) -> Set[int]: + layer = self.get_layer_by_id(layer_id) + neuron_count = self._neuron_count(layer, "out_neurons") + if neuron_count is None: + raise ValueError(f"Layer {layer_id} does not expose output neurons.") + + if neuron_indices is None: + indices = set() + elif isinstance(neuron_indices, int): + indices = {neuron_indices} + else: + indices = set(neuron_indices) + if not all(isinstance(index, int) for index in indices): + raise TypeError("neuron_indices must contain integers only.") + if not indices and not allow_empty: + raise ValueError("neuron_indices cannot be empty for this operation.") + + invalid = sorted( + index for index in indices if index < -neuron_count or index >= neuron_count + ) + if invalid: + raise ValueError( + f"Neuron indices {invalid} are outside layer {layer_id}'s " + f"valid range [-{neuron_count}, {neuron_count - 1}]." + ) + return indices + + def add_neurons(self, layer_id: int, count: int = 1) -> None: + """Append ``count`` neurons and propagate the change through the graph.""" + self.get_layer_by_id(layer_id) + if not isinstance(count, int) or isinstance(count, bool) or count < 1: + raise ValueError("count must be a positive integer.") + self.operate( + layer_id=layer_id, + neuron_indices=set(range(count)), + op_type=ArchitectureNeuronsOpType.ADD, + ) + + def prune_neurons(self, layer_id: int, neuron_indices: Iterable[int] | int) -> None: + """Remove selected neurons and update dependent layers.""" + indices = self._validated_neuron_indices( + layer_id, neuron_indices, allow_empty=False + ) + self.operate( + layer_id=layer_id, + neuron_indices=indices, + op_type=ArchitectureNeuronsOpType.PRUNE, + ) + + def freeze_neurons( + self, + layer_id: int, + neuron_indices: Optional[Iterable[int] | int] = None, + ) -> None: + """Freeze selected neurons idempotently (all when omitted).""" + self._set_neurons_frozen(layer_id, neuron_indices, frozen=True) + + def unfreeze_neurons( + self, + layer_id: int, + neuron_indices: Optional[Iterable[int] | int] = None, + ) -> None: + """Unfreeze selected neurons idempotently (all when omitted).""" + self._set_neurons_frozen(layer_id, neuron_indices, frozen=False) + + def _set_neurons_frozen( + self, + layer_id: int, + neuron_indices: Optional[Iterable[int] | int], + *, + frozen: bool, + ) -> None: + indices = self._validated_neuron_indices( + layer_id, neuron_indices, allow_empty=True + ) + layer = self.get_layer_by_id(layer_id) + weight = getattr(layer, "weight", None) + if weight is None: + raise ValueError( + f"Layer {layer_id} has no learnable weights to freeze. " + "Select a learnable layer from get_model_graph()." + ) + lr_overrides = getattr(layer, "neuron_2_lr", {}) + if not weight.requires_grad or "weight" not in lr_overrides: + raise ValueError( + f"Layer {layer_id} does not have a trainable, per-neuron " + "tracked weight. Freeze or unfreeze it through PyTorch before " + "wrapping the model." + ) + neuron_count = self._neuron_count(layer, "out_neurons") or 0 + selected = ( + {index % neuron_count for index in indices} + if indices + else set(range(neuron_count)) + ) + weight_lrs = lr_overrides["weight"] + currently_frozen = { + index for index in selected if float(weight_lrs.get(index, 1.0)) == 0.0 + } + indices_to_toggle = ( + selected - currently_frozen if frozen else currently_frozen + ) + if not indices_to_toggle: + return + self.operate( + layer_id=layer_id, + neuron_indices=indices_to_toggle, + op_type=ArchitectureNeuronsOpType.FREEZE, + ) + + def reset_neurons( + self, + layer_id: int, + neuron_indices: Optional[Iterable[int] | int] = None, + ) -> None: + """Reinitialize selected neurons (all when omitted).""" + indices = self._validated_neuron_indices( + layer_id, neuron_indices, allow_empty=True + ) + self.operate( + layer_id=layer_id, + neuron_indices=indices, + op_type=ArchitectureNeuronsOpType.RESET, + ) + + def perturb_neurons( + self, + layer_id: int, + neuron_indices: Optional[Iterable[int] | int] = None, + *, + ratio: float = 0.1, + ) -> None: + """Perturb selected weights by a relative random amount.""" + if not isinstance(ratio, (int, float)) or isinstance(ratio, bool) or not 0 < ratio < 1: + raise ValueError("ratio must be strictly between 0 and 1.") + indices = self._validated_neuron_indices( + layer_id, neuron_indices, allow_empty=True + ) + self.operate( + layer_id=layer_id, + neuron_indices=indices, + op_type=ArchitectureNeuronsOpType.RESET, + perturbation_ratio=float(ratio), + ) def register_dependencies( self,