Download src/pcm/planner/pythia_split_translate.py from Dermitio/Planner-Cache: direct link, hf CLI and curl.
- Browser
- Download file 8.02 kB
-
https://huggingface.co/Dermitio/Planner-Cache/resolve/main/src/pcm/planner/pythia_split_translate.py
- Command line
-
hf download hf://Dermitio/Planner-Cache/src/pcm/planner/pythia_split_translate.py
-
curl -L -o pythia_split_translate.py https://huggingface.co/Dermitio/Planner-Cache/resolve/main/src/pcm/planner/pythia_split_translate.py
8.02 kB
| """Frozen-Pythia attachment for the split canonical router/translator.""" | |
| from __future__ import annotations | |
| import torch | |
| from torch import Tensor, nn | |
| from pathlib import Path | |
| from pcm.planner.canonical import CanonicalPStore | |
| from pcm.planner.split_translator import ( | |
| ByteEntityEncoder, | |
| CanonicalPRouter, | |
| CanonicalRouterIndex, | |
| RouteResult, | |
| FactorizedCanonicalQuery, | |
| SplitPTranslatePackage, | |
| config_checksum, | |
| ) | |
| def pythia_model_identifier(base_model: nn.Module) -> str: | |
| configured = str(getattr(base_model.config, "_name_or_path", "")).strip() | |
| if configured and Path(configured).is_absolute(): | |
| configured = Path(configured).name | |
| return configured or str(getattr(base_model.config, "model_type", "gpt_neox")) | |
| class PythiaSplitTranslatedModel(nn.Module): | |
| def __init__( | |
| self, | |
| base_model: nn.Module, | |
| package: SplitPTranslatePackage, | |
| router: CanonicalPRouter, | |
| entity_encoder: ByteEntityEncoder, | |
| ) -> None: | |
| super().__init__() | |
| if not hasattr(base_model, "gpt_neox"): | |
| raise TypeError("base model must expose GPT-NeoX transformer layers") | |
| package.validate_compatibility( | |
| model_id=pythia_model_identifier(base_model), | |
| model_hidden_width=int(base_model.config.hidden_size), | |
| attachment_layers=package.config.attachment_layers, | |
| model_config_sha256=config_checksum(base_model.config), | |
| ) | |
| self.base_model = base_model | |
| self.package = package | |
| self.router = router | |
| self.entity_encoder = entity_encoder | |
| for parameter in base_model.parameters(): | |
| parameter.requires_grad_(False) | |
| layers = base_model.gpt_neox.layers | |
| if any(index < 0 or index >= len(layers) for index in package.config.attachment_layers): | |
| raise IndexError("split translator attachment layer is outside Pythia depth") | |
| self._store: CanonicalPStore | None = None | |
| self._index: CanonicalRouterIndex | None = None | |
| self._oracle_indices: Tensor | None = None | |
| self._query_entity_anchor: Tensor | None = None | |
| self._gate_enabled = True | |
| self._injection_enabled = True | |
| self._collect = False | |
| self._gate_telemetry: list[Tensor] = [] | |
| self._route_telemetry: list[RouteResult] = [] | |
| self._query_telemetry: list[FactorizedCanonicalQuery] = [] | |
| self._handles = [ | |
| layers[index].register_forward_hook(self._hook) | |
| for index in package.config.attachment_layers | |
| ] | |
| self.base_model.eval() | |
| def _oracle_route(self, query, hidden: Tensor) -> RouteResult: | |
| assert self._index is not None and self._oracle_indices is not None | |
| scores, features = self.router.all_scores(query, self._index) | |
| batch, sequence = hidden.shape[:2] | |
| indices = self._oracle_indices.to(hidden.device).view(batch, 1, 1).expand(batch, sequence, 1) | |
| selected_scores = scores.gather(-1, indices) | |
| selected_features = features.gather( | |
| -2, indices.unsqueeze(-1).expand(batch, sequence, 1, 4) | |
| ) | |
| return RouteResult( | |
| indices=indices, | |
| scores=selected_scores, | |
| weights=torch.ones_like(selected_scores), | |
| features=selected_features, | |
| accepted=torch.ones((batch, sequence), dtype=torch.bool, device=hidden.device), | |
| has_valid=True, | |
| ) | |
| def _hook(self, _module, _inputs, hidden: Tensor): | |
| if self._store is None or self._store.cache.occupied == 0: | |
| return hidden | |
| assert self._index is not None | |
| entity_anchor = None | |
| if self._query_entity_anchor is not None: | |
| entity_anchor = self._query_entity_anchor[:, None, :].expand( | |
| hidden.shape[0], hidden.shape[1], -1 | |
| ) | |
| query = self.package.query_projector(hidden, entity_anchor=entity_anchor) | |
| if self._collect: | |
| self._query_telemetry.append(FactorizedCanonicalQuery( | |
| entity=query.entity.detach(), | |
| relation_logits=query.relation_logits.detach(), | |
| metadata_logits=query.metadata_logits.detach(), | |
| )) | |
| route = ( | |
| self._oracle_route(query, hidden) | |
| if self._oracle_indices is not None | |
| else self.router.route(query, self._index, top_k=self.package.config.top_k) | |
| ) | |
| if self._collect: | |
| self._route_telemetry.append(RouteResult( | |
| indices=route.indices.detach(), scores=route.scores.detach(), | |
| weights=route.weights.detach(), features=route.features.detach(), | |
| accepted=route.accepted.detach(), | |
| has_valid=route.has_valid, | |
| )) | |
| if not self._injection_enabled or not route.has_valid: | |
| return hidden | |
| canonical = self._store.canonical_values.to( | |
| device=hidden.device, dtype=route.weights.dtype | |
| ) | |
| selected = canonical[route.indices] | |
| pooled = torch.einsum("...k,...kd->...d", route.weights, selected) | |
| translated = self.package.value_translator(pooled) | |
| route_features = torch.einsum( | |
| "...k,...kf->...f", route.weights, route.features | |
| ) | |
| gate = ( | |
| self.package.gate(hidden, translated, route_features) | |
| if self._gate_enabled | |
| else torch.ones(hidden.shape[:-1], device=hidden.device, dtype=translated.dtype) | |
| ) | |
| gate = gate * route.accepted.to(gate.dtype) | |
| if self._collect: | |
| self._gate_telemetry.append(gate.detach()) | |
| return hidden + (gate.unsqueeze(-1) * translated).to(hidden.dtype) | |
| def train(self, mode: bool = True): | |
| super().train(mode) | |
| self.base_model.eval() | |
| self.package.train(mode) | |
| self.router.train(mode) | |
| return self | |
| def forward( | |
| self, | |
| *args, | |
| p_store: CanonicalPStore | None = None, | |
| query_entity_surfaces: list[str] | tuple[str, ...] | None = None, | |
| oracle_indices: Tensor | None = None, | |
| gate_enabled: bool = True, | |
| injection_enabled: bool = True, | |
| collect_telemetry: bool = False, | |
| **kwargs, | |
| ): | |
| if self._store is not None: | |
| raise RuntimeError("PythiaSplitTranslatedModel is not reentrant") | |
| self._store = p_store | |
| self._oracle_indices = oracle_indices | |
| if query_entity_surfaces is not None: | |
| if "input_ids" in kwargs and len(query_entity_surfaces) != kwargs["input_ids"].shape[0]: | |
| raise ValueError("query entity surface count must match the input batch") | |
| self._query_entity_anchor = self.entity_encoder( | |
| list(query_entity_surfaces) | |
| ).to(next(self.package.parameters()).device) | |
| else: | |
| self._query_entity_anchor = None | |
| self._gate_enabled = gate_enabled | |
| self._injection_enabled = injection_enabled | |
| self._collect = collect_telemetry | |
| self._gate_telemetry.clear() | |
| self._route_telemetry.clear() | |
| self._query_telemetry.clear() | |
| if p_store is not None and p_store.cache.occupied: | |
| self._index = self.router.build_index( | |
| p_store, self.entity_encoder, device=next(self.package.parameters()).device | |
| ) | |
| try: | |
| return self.base_model(*args, **kwargs) | |
| finally: | |
| self._store = None | |
| self._index = None | |
| self._oracle_indices = None | |
| self._query_entity_anchor = None | |
| self._collect = False | |
| def gate_telemetry(self): | |
| return tuple(self._gate_telemetry) | |
| def route_telemetry(self): | |
| return tuple(self._route_telemetry) | |
| def query_telemetry(self): | |
| return tuple(self._query_telemetry) | |
| def close(self): | |
| for handle in self._handles: | |
| handle.remove() | |
| self._handles.clear() | |