| 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())) |
|
|