File size: 5,445 Bytes
7ee20db | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 | 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()))
|