from __future__ import annotations from dataclasses import dataclass, field import json from pathlib import Path from typing import Any @dataclass class Order2Result: """Measurements and selections produced by the order-2 interaction method.""" depth: int baseline_nll: float single_nll: dict[int, float] pair_nll: dict[tuple[int, int], float] first_order: dict[int, float] = field(default_factory=dict) second_order: dict[tuple[int, int], float] = field(default_factory=dict) delete_order: list[int] = field(default_factory=list) greedy_path: list[dict[str, Any]] = field(default_factory=list) def build_interactions(self) -> "Order2Result": d0 = float(self.baseline_nll) self.first_order = { i: float(self.single_nll[i]) - d0 for i in range(self.depth) } self.second_order = {} for i in range(self.depth): for j in range(i + 1, self.depth): self.second_order[(i, j)] = ( float(self.pair_nll[(i, j)]) - float(self.single_nll[i]) - float(self.single_nll[j]) + d0 ) return self def build_greedy_path(self, max_delete: int | None = None) -> "Order2Result": if not self.first_order or not self.second_order: self.build_interactions() if max_delete is None: max_delete = self.depth - 1 if not 0 <= max_delete < self.depth: raise ValueError(f"max_delete must be in [0, {self.depth - 1}]") deleted: list[int] = [] deleted_set: set[int] = set() path: list[dict[str, Any]] = [] cumulative = 0.0 for step in range(max_delete): candidates: list[tuple[float, int]] = [] for i in range(self.depth): if i in deleted_set: continue interaction = sum( self.second_order[tuple(sorted((i, j)))] for j in deleted ) marginal = self.first_order[i] + interaction candidates.append((marginal, i)) marginal, chosen = min(candidates, key=lambda z: (z[0], z[1])) deleted.append(chosen) deleted_set.add(chosen) cumulative += marginal path.append( { "step": step + 1, "deleted_layer": chosen, "marginal_predicted_nll_change": float(marginal), "cumulative_predicted_nll_change": float(cumulative), } ) self.delete_order = deleted self.greedy_path = path return self def select(self, target_layers: int) -> dict[str, list[int]]: if not 1 <= target_layers <= self.depth: raise ValueError(f"target_layers must be in [1, {self.depth}]") n_delete = self.depth - target_layers if len(self.delete_order) < n_delete: self.build_greedy_path(max_delete=n_delete) deleted = list(self.delete_order[:n_delete]) deleted_set = set(deleted) retained = [i for i in range(self.depth) if i not in deleted_set] return {"retained_layers": retained, "deleted_layers": deleted} def to_dict(self) -> dict[str, Any]: return { "method": "order-2 interaction greedy", "depth": self.depth, "baseline_nll": float(self.baseline_nll), "single_nll": {str(k): float(v) for k, v in self.single_nll.items()}, "pair_nll": {f"{i},{j}": float(v) for (i, j), v in self.pair_nll.items()}, "first_order_delta": {str(k): float(v) for k, v in self.first_order.items()}, "second_order_interaction": { f"{i},{j}": float(v) for (i, j), v in self.second_order.items() }, "delete_order": list(self.delete_order), "greedy_path": list(self.greedy_path), "complete": True, } def save_json(self, path: str | Path) -> None: path = Path(path) path.parent.mkdir(parents=True, exist_ok=True) tmp = path.with_suffix(path.suffix + ".tmp") tmp.write_text(json.dumps(self.to_dict(), indent=2)) tmp.replace(path) @classmethod def from_dict(cls, obj: dict[str, Any]) -> "Order2Result": pair = {} for key, value in obj.get("pair_nll", {}).items(): i, j = (int(x) for x in key.split(",")) pair[(i, j)] = float(value) second = {} for key, value in obj.get("second_order_interaction", {}).items(): i, j = (int(x) for x in key.split(",")) second[(i, j)] = float(value) result = cls( depth=int(obj["depth"]), baseline_nll=float(obj["baseline_nll"]), single_nll={int(k): float(v) for k, v in obj.get("single_nll", {}).items()}, pair_nll=pair, first_order={ int(k): float(v) for k, v in obj.get("first_order_delta", {}).items() }, second_order=second, delete_order=[int(x) for x in obj.get("delete_order", [])], greedy_path=list(obj.get("greedy_path", [])), ) return result @classmethod def load_json(cls, path: str | Path) -> "Order2Result": return cls.from_dict(json.loads(Path(path).read_text()))