Something went wrong. Try again.
This repository has no description
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124#!/usr/bin/env python3"""Does each member of the colony actually matter?
Silences units at inference and measures what the colony loses. Two outcomes,opposite implications:
- Removing a unit costs little -> members are redundant echoes of the shared channel; the colony is one voice with extra steps. - Removing a unit costs a lot -> each member's private evidence is load-bearing; the group depends on a genuine division of epistemic labour.
For a distributed-evidence colony this is the sharpest available test, because asilenced unit takes its private clues out of the discussion entirely.
Usage: python scripts/colony_ablate.py --checkpoint <ckpt> --config <all_config.yaml> \ --data data/sudoku-testsub-12k --batches 3"""
import argparseimport itertoolsimport jsonimport osimport sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
import torch # noqa: E402import yaml # noqa: E402
from dsem.device import get_device # noqa: E402from dsem.pretrain import PretrainConfig, create_dataloader, create_model # noqa: E402
IGNORE = -100
def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--checkpoint", required=True) ap.add_argument("--config", required=True) ap.add_argument("--data", required=True) ap.add_argument("--batches", type=int, default=3) ap.add_argument("--out", default=None) args = ap.parse_args()
device = get_device() with open(args.config) as f: cfg_dict = yaml.unsafe_load(f) config = PretrainConfig(**cfg_dict) config.data_paths_test = [args.data] config.evaluators = []
loader, metadata = create_dataloader(config, "test", test_set_mode=True, epochs_per_iter=1, global_batch_size=config.global_batch_size, rank=0, world_size=1) model, _, _ = create_model(config, metadata, rank=0, world_size=1) model.load_state_dict(torch.load(args.checkpoint, map_location=device), assign=True) model.eval()
inner = model.model.inner if hasattr(model, "model") else model.inner if hasattr(inner, "_orig_mod"): inner = inner._orig_mod K = inner.config.num_agents
from dsem.models.recursive_reasoning.colony import ColonyInnerCarry
# which subsets of units get to speak subsets = [("all K units", list(range(K)))] if K > 1: subsets.append((f"drop 1 unit ({K-1} speak)", list(range(K - 1)))) if K > 2: subsets.append((f"drop half ({K//2} speak)", list(range(K // 2)))) subsets.append(("1 unit only", [0]))
batches = [] for bi, (_set, batch, _gbs) in enumerate(loader): if bi >= args.batches: break batches.append({k: v.to(device) for k, v in batch.items()})
results = [] print(f"colony K={K}, {sum(b['labels'].shape[0] for b in batches)} puzzles\n") print(f"{'configuration':>24} {'solved':>9} {'vs full':>9}") full = None for name, keep in subsets: active = torch.tensor(keep, device=device, dtype=torch.long) solved_total, seen = 0.0, 0 for batch in batches: labels = batch["labels"] mask = labels != IGNORE B = labels.shape[0] with torch.device(device): carry = model.initial_carry(batch) ic = inner.reset_carry(torch.ones(B, dtype=torch.bool, device=device), carry.inner_carry) for _ in range(inner.config.halt_max_steps): out = inner.forward_instrumented(ic, batch, active=active) ic = ColonyInnerCarry(y=out["y"], z=out["z"]) pred = out["rounds"][-1]["aggregate"] solved_total += float(((pred == labels) | ~mask).all(-1).sum()) seen += B acc = solved_total / max(seen, 1) * 100 if full is None: full = acc delta = acc - full results.append({"config": name, "units_speaking": len(keep), "exact_accuracy": acc, "delta": delta}) print(f"{name:>24} {acc:>8.2f}% {delta:>+8.2f}")
loss_one = results[1]["delta"] if len(results) > 1 else 0.0 print("\ndiagnosis:") if abs(loss_one) < 0.5: print(" REDUNDANT — silencing a member costs almost nothing. The units are") print(" interchangeable; the colony is not dividing epistemic labour.") else: print(f" LOAD-BEARING — silencing one member costs {abs(loss_one):.1f} points.") print(" Each unit contributes evidence the others cannot supply.")
if args.out: os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True) with open(args.out, "w") as f: json.dump({"K": K, "results": results}, f, indent=2)
if __name__ == "__main__": main()