Download layer_interactions/result.py from Angshul/LayerInteractions: direct link, hf CLI and curl.
- Browser
- Download file 5.45 kB
-
https://huggingface.co/Angshul/LayerInteractions/resolve/main/layer_interactions/result.py
- Command line
-
hf download hf://Angshul/LayerInteractions/layer_interactions/result.py
-
curl -L -o result.py https://huggingface.co/Angshul/LayerInteractions/resolve/main/layer_interactions/result.py
5.45 kB
| from __future__ import annotations | |
| from dataclasses import dataclass, field | |
| import json | |
| from pathlib import Path | |
| from typing import Any | |
| 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) | |
| 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 | |
| def load_json(cls, path: str | Path) -> "Order2Result": | |
| return cls.from_dict(json.loads(Path(path).read_text())) | |