Download layer_interactions/pruner.py from Angshul/LayerInteractions: direct link, hf CLI and curl.
- Browser
- Download file 16.8 kB
-
https://huggingface.co/Angshul/LayerInteractions/resolve/main/layer_interactions/pruner.py
- Command line
-
hf download hf://Angshul/LayerInteractions/layer_interactions/pruner.py
-
curl -L -o pruner.py https://huggingface.co/Angshul/LayerInteractions/resolve/main/layer_interactions/pruner.py
16.8 kB
| from __future__ import annotations | |
| import gc | |
| import itertools | |
| import json | |
| import os | |
| from pathlib import Path | |
| import random | |
| import shutil | |
| import tempfile | |
| from typing import Iterable, Sequence | |
| import psutil | |
| import torch | |
| import torch.nn.functional as F | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, PreTrainedModel, PreTrainedTokenizerBase | |
| from .result import Order2Result | |
| _LAYER_PATHS = ( | |
| "model.layers", # Llama/Mistral/DeepSeek/OLMoE and compatible models | |
| "transformer.h", # GPT-2 style | |
| "gpt_neox.layers", # GPT-NeoX style | |
| ) | |
| def _resolve_dtype(dtype: str | torch.dtype) -> torch.dtype: | |
| if isinstance(dtype, torch.dtype): | |
| return dtype | |
| table = { | |
| "bfloat16": torch.bfloat16, | |
| "bf16": torch.bfloat16, | |
| "float16": torch.float16, | |
| "fp16": torch.float16, | |
| "float32": torch.float32, | |
| "fp32": torch.float32, | |
| } | |
| try: | |
| return table[dtype.lower()] | |
| except KeyError as exc: | |
| raise ValueError(f"Unsupported dtype: {dtype}") from exc | |
| def _get_attr_path(obj, path: str): | |
| for part in path.split("."): | |
| obj = getattr(obj, part) | |
| return obj | |
| def _set_attr_path(obj, path: str, value) -> None: | |
| parts = path.split(".") | |
| parent = obj | |
| for part in parts[:-1]: | |
| parent = getattr(parent, part) | |
| setattr(parent, parts[-1], value) | |
| class Order2Pruner: | |
| """Second-order interaction layer pruning for Hugging Face causal LMs. | |
| The method measures the NLL after deleting each layer and each pair of | |
| layers, forms a quadratic deletion-loss surrogate, then greedily deletes | |
| the layer with minimum current marginal predicted NLL increase. | |
| """ | |
| def __init__( | |
| self, | |
| model: str | Path | PreTrainedModel, | |
| tokenizer: str | Path | PreTrainedTokenizerBase | None = None, | |
| *, | |
| dtype: str | torch.dtype = "bfloat16", | |
| layer_path: str | None = None, | |
| device_map: str | dict | None = "auto", | |
| gpu_memory_gib: int | None = None, | |
| cpu_memory_gib: int | None = None, | |
| activation_headroom_gib: float = 4.0, | |
| offload_dir: str | Path | None = None, | |
| local_files_only: bool = False, | |
| trust_remote_code: bool = False, | |
| ) -> None: | |
| self.model_source = model | |
| self.tokenizer_source = tokenizer | |
| self.dtype = _resolve_dtype(dtype) | |
| self.layer_path = layer_path | |
| self.device_map = device_map | |
| self.gpu_memory_gib = gpu_memory_gib | |
| self.cpu_memory_gib = cpu_memory_gib | |
| self.activation_headroom_gib = float(activation_headroom_gib) | |
| self.local_files_only = local_files_only | |
| self.trust_remote_code = trust_remote_code | |
| self.model: PreTrainedModel | None = model if isinstance(model, PreTrainedModel) else None | |
| self.tokenizer: PreTrainedTokenizerBase | None = ( | |
| tokenizer if isinstance(tokenizer, PreTrainedTokenizerBase) else None | |
| ) | |
| self._owns_offload_dir = offload_dir is None | |
| self.offload_dir = Path(offload_dir) if offload_dir else Path( | |
| tempfile.mkdtemp(prefix="layer_interactions_offload_") | |
| ) | |
| self.offload_dir.mkdir(parents=True, exist_ok=True) | |
| self._resolved_layer_path: str | None = None | |
| self.result: Order2Result | None = None | |
| def depth(self) -> int: | |
| self._ensure_loaded() | |
| return len(self._layers()) | |
| def _auto_max_memory(self) -> dict | None: | |
| if not torch.cuda.is_available(): | |
| return None | |
| free_gib = torch.cuda.mem_get_info()[0] / 2**30 | |
| gpu = self.gpu_memory_gib | |
| if gpu is None: | |
| gpu = max(1, int(free_gib - self.activation_headroom_gib)) | |
| cpu_free_gib = psutil.virtual_memory().available / 2**30 | |
| cpu = self.cpu_memory_gib | |
| if cpu is None: | |
| cpu = max(2, int(cpu_free_gib - 4.0)) | |
| return {0: f"{gpu}GiB", "cpu": f"{cpu}GiB"} | |
| def _ensure_loaded(self) -> None: | |
| if self.tokenizer is None: | |
| source = self.tokenizer_source or self.model_source | |
| if isinstance(source, PreTrainedModel): | |
| raise ValueError("Pass a tokenizer when model is an already-instantiated model.") | |
| self.tokenizer = AutoTokenizer.from_pretrained( | |
| source, | |
| local_files_only=self.local_files_only, | |
| trust_remote_code=self.trust_remote_code, | |
| use_fast=True, | |
| ) | |
| if self.tokenizer.pad_token_id is None: | |
| self.tokenizer.pad_token = self.tokenizer.eos_token | |
| if self.model is None: | |
| kwargs = dict( | |
| local_files_only=self.local_files_only, | |
| trust_remote_code=self.trust_remote_code, | |
| dtype=self.dtype, | |
| low_cpu_mem_usage=True, | |
| ) | |
| if self.device_map is not None: | |
| kwargs["device_map"] = self.device_map | |
| max_memory = self._auto_max_memory() | |
| if max_memory is not None and self.device_map == "auto": | |
| kwargs.update( | |
| max_memory=max_memory, | |
| offload_folder=str(self.offload_dir), | |
| offload_state_dict=True, | |
| offload_buffers=True, | |
| ) | |
| self.model = AutoModelForCausalLM.from_pretrained(self.model_source, **kwargs) | |
| self.model.config.use_cache = False | |
| if hasattr(self.model, "generation_config"): | |
| self.model.generation_config.use_cache = False | |
| self.model.eval() | |
| self._resolve_layer_path() | |
| def _resolve_layer_path(self) -> str: | |
| if self._resolved_layer_path is not None: | |
| return self._resolved_layer_path | |
| self._ensure_model_exists_for_resolution() | |
| candidates = (self.layer_path,) if self.layer_path else _LAYER_PATHS | |
| for path in candidates: | |
| if path is None: | |
| continue | |
| try: | |
| value = _get_attr_path(self.model, path) | |
| except AttributeError: | |
| continue | |
| if isinstance(value, (torch.nn.ModuleList, list, tuple)): | |
| self._resolved_layer_path = path | |
| return path | |
| raise ValueError( | |
| "Could not find transformer layers automatically. " | |
| "Pass layer_path, e.g. layer_path='model.layers'." | |
| ) | |
| def _ensure_model_exists_for_resolution(self) -> None: | |
| if self.model is None: | |
| raise RuntimeError("Model has not been loaded.") | |
| def _layers(self): | |
| self._ensure_model_exists_for_resolution() | |
| path = self._resolved_layer_path or self._resolve_layer_path() | |
| return _get_attr_path(self.model, path) | |
| def _set_layers(self, layers: Sequence[torch.nn.Module]) -> None: | |
| self._ensure_model_exists_for_resolution() | |
| path = self._resolved_layer_path or self._resolve_layer_path() | |
| _set_attr_path(self.model, path, torch.nn.ModuleList(list(layers))) | |
| for k, block in enumerate(self._layers()): | |
| if hasattr(block, "layer_idx"): | |
| block.layer_idx = k | |
| if hasattr(block, "self_attn") and hasattr(block.self_attn, "layer_idx"): | |
| block.self_attn.layer_idx = k | |
| if hasattr(self.model.config, "num_hidden_layers"): | |
| self.model.config.num_hidden_layers = len(layers) | |
| def calibration_batches( | |
| self, | |
| texts: Sequence[str], | |
| *, | |
| n_sequences: int = 32, | |
| sequence_length: int = 128, | |
| seed: int = 42, | |
| tokenizer_chunk_size: int = 4096, | |
| ) -> list[torch.Tensor]: | |
| self._ensure_loaded() | |
| joined = "\n\n".join(str(x) for x in texts if str(x).strip()) | |
| enc = self.tokenizer( | |
| joined, | |
| add_special_tokens=False, | |
| truncation=True, | |
| max_length=tokenizer_chunk_size, | |
| return_overflowing_tokens=True, | |
| return_attention_mask=False, | |
| ) | |
| ids = torch.tensor( | |
| list(itertools.chain.from_iterable(enc["input_ids"])), dtype=torch.long | |
| ) | |
| max_start = len(ids) - sequence_length - 1 | |
| if max_start < 0: | |
| raise ValueError( | |
| f"Calibration corpus has {len(ids)} tokens; need at least {sequence_length + 1}." | |
| ) | |
| if max_start + 1 < n_sequences: | |
| raise ValueError( | |
| f"Not enough distinct start positions for {n_sequences} calibration sequences." | |
| ) | |
| starts = random.Random(seed).sample(range(max_start + 1), n_sequences) | |
| return [ids[s : s + sequence_length].unsqueeze(0) for s in starts] | |
| def _input_device(self) -> torch.device: | |
| emb = self.model.get_input_embeddings() | |
| for p in emb.parameters(): | |
| if p.device.type != "meta": | |
| return p.device | |
| return torch.device("cpu") | |
| def score_nll(self, batches: Iterable[torch.Tensor]) -> float: | |
| self._ensure_loaded() | |
| dev = self._input_device() | |
| total = 0.0 | |
| ntok = 0 | |
| with torch.inference_mode(): | |
| for cpu_x in batches: | |
| x = cpu_x.to(dev) | |
| output = self.model(input_ids=x, use_cache=False) | |
| logits = output.logits[:, :-1, :] | |
| target = x[:, 1:].to(logits.device) | |
| loss = F.cross_entropy( | |
| logits.reshape(-1, logits.shape[-1]).float(), | |
| target.reshape(-1), | |
| reduction="sum", | |
| ) | |
| total += float(loss.detach().cpu()) | |
| ntok += int(target.numel()) | |
| del x, output, logits, target, loss | |
| if ntok == 0: | |
| raise ValueError("No calibration tokens were scored.") | |
| return total / ntok | |
| def _atomic_write(path: Path, obj: dict) -> None: | |
| path.parent.mkdir(parents=True, exist_ok=True) | |
| tmp = path.with_suffix(path.suffix + ".tmp") | |
| tmp.write_text(json.dumps(obj, indent=2)) | |
| os.replace(tmp, path) | |
| def fit( | |
| self, | |
| *, | |
| texts: Sequence[str] | None = None, | |
| batches: Sequence[torch.Tensor] | None = None, | |
| n_sequences: int = 32, | |
| sequence_length: int = 128, | |
| seed: int = 42, | |
| checkpoint_path: str | Path | None = None, | |
| resume: bool = True, | |
| max_delete: int | None = None, | |
| ) -> Order2Result: | |
| """Measure baseline/single/pair NLLs and construct the greedy order-2 path.""" | |
| self._ensure_loaded() | |
| if batches is None: | |
| if texts is None: | |
| raise ValueError("Pass either texts=... or batches=...") | |
| batches = self.calibration_batches( | |
| texts, | |
| n_sequences=n_sequences, | |
| sequence_length=sequence_length, | |
| seed=seed, | |
| ) | |
| depth = self.depth | |
| ckpt = Path(checkpoint_path) if checkpoint_path else None | |
| state = { | |
| "method": "order-2 interaction greedy", | |
| "depth": depth, | |
| "baseline_nll": None, | |
| "single_nll": {}, | |
| "pair_nll": {}, | |
| "complete": False, | |
| } | |
| if ckpt and resume and ckpt.exists(): | |
| loaded = json.loads(ckpt.read_text()) | |
| if int(loaded.get("depth", depth)) != depth: | |
| raise ValueError("Checkpoint depth does not match the loaded model.") | |
| state.update(loaded) | |
| state.setdefault("single_nll", {}) | |
| state.setdefault("pair_nll", {}) | |
| original = list(self._layers()) | |
| try: | |
| if state["baseline_nll"] is None: | |
| self._set_layers(original) | |
| state["baseline_nll"] = self.score_nll(batches) | |
| if ckpt: | |
| self._atomic_write(ckpt, state) | |
| for i in range(depth): | |
| key = str(i) | |
| if key in state["single_nll"]: | |
| continue | |
| kept = [b for j, b in enumerate(original) if j != i] | |
| self._set_layers(kept) | |
| state["single_nll"][key] = self.score_nll(batches) | |
| if ckpt: | |
| self._atomic_write(ckpt, state) | |
| self._set_layers(original) | |
| total_pairs = depth * (depth - 1) // 2 | |
| for i in range(depth): | |
| for j in range(i + 1, depth): | |
| key = f"{i},{j}" | |
| if key in state["pair_nll"]: | |
| continue | |
| kept = [b for q, b in enumerate(original) if q not in (i, j)] | |
| self._set_layers(kept) | |
| state["pair_nll"][key] = self.score_nll(batches) | |
| if ckpt: | |
| self._atomic_write(ckpt, state) | |
| self._set_layers(original) | |
| print(f"pair {key:>7s} | {len(state['pair_nll'])}/{total_pairs}", flush=True) | |
| result = Order2Result( | |
| depth=depth, | |
| baseline_nll=float(state["baseline_nll"]), | |
| single_nll={int(k): float(v) for k, v in state["single_nll"].items()}, | |
| pair_nll={ | |
| tuple(int(x) for x in k.split(",")): float(v) | |
| for k, v in state["pair_nll"].items() | |
| }, | |
| ) | |
| result.build_interactions().build_greedy_path(max_delete=max_delete) | |
| self.result = result | |
| if ckpt: | |
| result.save_json(ckpt) | |
| return result | |
| finally: | |
| self._set_layers(original) | |
| def select(self, target_layers: int, result: Order2Result | None = None) -> dict[str, list[int]]: | |
| result = result or self.result | |
| if result is None: | |
| raise RuntimeError("Call fit() first or pass result=...") | |
| return result.select(target_layers) | |
| def apply( | |
| self, | |
| target_layers: int, | |
| *, | |
| result: Order2Result | None = None, | |
| ) -> PreTrainedModel: | |
| """Apply a selection to the currently loaded model in place.""" | |
| self._ensure_loaded() | |
| result = result or self.result | |
| if result is None: | |
| raise RuntimeError("Call fit() first or pass result=...") | |
| selection = result.select(target_layers) | |
| original = list(self._layers()) | |
| retained = selection["retained_layers"] | |
| self._set_layers([original[i] for i in retained]) | |
| return self.model | |
| def prune( | |
| self, | |
| target_layers: int, | |
| *, | |
| texts: Sequence[str] | None = None, | |
| batches: Sequence[torch.Tensor] | None = None, | |
| n_sequences: int = 32, | |
| sequence_length: int = 128, | |
| seed: int = 42, | |
| checkpoint_path: str | Path | None = None, | |
| resume: bool = True, | |
| ) -> tuple[PreTrainedModel, Order2Result]: | |
| result = self.fit( | |
| texts=texts, | |
| batches=batches, | |
| n_sequences=n_sequences, | |
| sequence_length=sequence_length, | |
| seed=seed, | |
| checkpoint_path=checkpoint_path, | |
| resume=resume, | |
| max_delete=self.depth - target_layers, | |
| ) | |
| model = self.apply(target_layers, result=result) | |
| return model, result | |
| def save_pruned( | |
| self, | |
| output_dir: str | Path, | |
| target_layers: int, | |
| *, | |
| result: Order2Result | None = None, | |
| safe_serialization: bool = True, | |
| ) -> Path: | |
| """Apply a selection and save the pruned model/tokenizer with HF save_pretrained().""" | |
| model = self.apply(target_layers, result=result) | |
| output_dir = Path(output_dir) | |
| output_dir.mkdir(parents=True, exist_ok=True) | |
| model.save_pretrained(output_dir, safe_serialization=safe_serialization) | |
| if self.tokenizer is not None: | |
| self.tokenizer.save_pretrained(output_dir) | |
| (output_dir / "layer_interaction_selection.json").write_text( | |
| json.dumps((result or self.result).select(target_layers), indent=2) | |
| ) | |
| return output_dir | |
| def close(self, *, drop_model: bool = False) -> None: | |
| if drop_model: | |
| self.model = None | |
| self.tokenizer = None | |
| gc.collect() | |
| if torch.cuda.is_available(): | |
| torch.cuda.empty_cache() | |
| try: | |
| torch.cuda.ipc_collect() | |
| except Exception: | |
| pass | |
| if self._owns_offload_dir: | |
| shutil.rmtree(self.offload_dir, ignore_errors=True) | |
| def __enter__(self) -> "Order2Pruner": | |
| return self | |
| def __exit__(self, exc_type, exc, tb) -> None: | |
| self.close(drop_model=True) | |