| """Shared checkpoint construction for the STRATA-COMPOSE replacement frontier.""" |
|
|
| from __future__ import annotations |
|
|
| import hashlib |
| from pathlib import Path |
|
|
| import torch |
|
|
| from strata.modeling.compose import ExportedAlgebraBlock, PrunableAlgebraLM |
| from strata.modeling.ph_core import DenseContinuationForCausalLM |
| from strata.modeling.ph_pat import PHPATConfig |
|
|
|
|
| def sha256_file(path: Path) -> str: |
| digest = hashlib.sha256() |
| with path.open("rb") as handle: |
| for chunk in iter(lambda: handle.read(16 << 20), b""): |
| digest.update(chunk) |
| return digest.hexdigest() |
|
|
|
|
| def _load_with_periodic_position_extension( |
| model: torch.nn.Module, |
| checkpoint_path: Path, |
| *, |
| position_name: str, |
| ) -> dict[str, object]: |
| """Load a checkpoint while periodically extending its sole position table.""" |
|
|
| source = torch.load(checkpoint_path, map_location="cpu", weights_only=True) |
| target = model.state_dict() |
| if set(source) != set(target): |
| missing = sorted(set(target) - set(source)) |
| extra = sorted(set(source) - set(target)) |
| raise ValueError(f"checkpoint keys differ: missing={missing[:3]}, extra={extra[:3]}") |
| source_positions = source[position_name] |
| target_positions = target[position_name] |
| if ( |
| source_positions.ndim != 2 |
| or target_positions.ndim != 2 |
| or source_positions.shape[1] != target_positions.shape[1] |
| or source_positions.shape[0] > target_positions.shape[0] |
| ): |
| raise ValueError("position embedding shapes cannot be extended") |
| state = {} |
| for name, value in source.items(): |
| if name == position_name: |
| positions = target_positions.clone() |
| source_length = source_positions.shape[0] |
| for start in range(0, target_positions.shape[0], source_length): |
| count = min(source_length, target_positions.shape[0] - start) |
| positions[start:start + count] = source_positions[:count] |
| state[name] = positions |
| else: |
| if value.shape != target[name].shape: |
| raise ValueError(f"non-position tensor changed shape: {name}") |
| state[name] = value |
| model.load_state_dict(state, strict=True) |
| return { |
| "source_context": source_positions.shape[0], |
| "target_context": target_positions.shape[0], |
| "extension": "periodic_copy", |
| "non_position_tensors_exact": True, |
| } |
|
|
|
|
| def load_dense_base( |
| config_path: Path, |
| checkpoint_path: Path, |
| device: torch.device, |
| *, |
| allow_position_extension: bool = False, |
| ) -> tuple[DenseContinuationForCausalLM, PHPATConfig]: |
| config = PHPATConfig.from_json_file(config_path).with_arm("dense_global") |
| model = DenseContinuationForCausalLM(config, gradient_checkpointing=False) |
| if allow_position_extension: |
| _load_with_periodic_position_extension( |
| model, checkpoint_path, position_name="position_embeddings.weight", |
| ) |
| else: |
| model.load_state_dict( |
| torch.load(checkpoint_path, map_location="cpu", weights_only=True), |
| strict=True, |
| ) |
| dtype = torch.bfloat16 if device.type == "cuda" else torch.float32 |
| return model.to(device=device, dtype=dtype), config |
|
|
|
|
| def build_exported_stage( |
| *, |
| config_path: Path, |
| base_checkpoint: Path, |
| stage_checkpoint: Path, |
| removed_layers: tuple[int, ...], |
| device: torch.device, |
| gamma_max: float = 0.02, |
| allow_position_extension: bool = False, |
| ) -> tuple[PrunableAlgebraLM, PHPATConfig]: |
| base, config = load_dense_base( |
| config_path, base_checkpoint, device, |
| allow_position_extension=allow_position_extension, |
| ) |
| model = PrunableAlgebraLM( |
| base, |
| config, |
| removed_layers, |
| gamma_max=gamma_max, |
| ) |
| model.export_layers(removed_layers) |
| if allow_position_extension: |
| _load_with_periodic_position_extension( |
| model, stage_checkpoint, |
| position_name="base_model.position_embeddings.weight", |
| ) |
| else: |
| model.load_state_dict( |
| torch.load(stage_checkpoint, map_location=device, weights_only=True), |
| strict=True, |
| ) |
| if model.physically_removed_layers() != tuple(sorted(removed_layers)): |
| raise AssertionError("stage checkpoint does not match its physical-removal manifest") |
| if model.dense_modules_in_replacement_layers(): |
| raise AssertionError("dense attention remains in a physically removed layer") |
| return model, config |
|
|
|
|
| def graph_adapter_state(model: PrunableAlgebraLM, layer: int) -> dict[str, torch.Tensor]: |
| block = model.base_model.blocks[layer] |
| if not isinstance(block, ExportedAlgebraBlock): |
| raise TypeError(f"graph-read layer {layer} is not physically exported") |
| return {name: value.detach().cpu() for name, value in block.graph_adapter.state_dict().items()} |
|
|
|
|
| __all__ = [ |
| "build_exported_stage", |
| "graph_adapter_state", |
| "load_dense_base", |
| "sha256_file", |
| ] |
|
|