Angshul's picture
Upload 6 files
7ee20db verified
Raw
History Blame Contribute Delete
5.45 kB
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()))