SWD / swd_loader.py
XuhanH's picture
Publish SWD factor checkpoints
43bcf0c verified
Raw
History Blame Contribute Delete
7.74 kB
"""Load SWD factor checkpoints into an already constructed PyTorch model."""
from __future__ import annotations
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Literal
import torch
from safetensors.torch import load_file
from torch import nn
@dataclass(frozen=True)
class AppliedReplacement:
module_path: str
input_dim: int
rank: int
output_dim: int
mode: str
class SWDLinear(nn.Module):
"""Two-factor linear map with explicit scalar bottleneck activations."""
def __init__(
self,
read: torch.Tensor,
write: torch.Tensor,
bias: torch.Tensor | None = None,
) -> None:
super().__init__()
if read.ndim != 2 or write.ndim != 2 or read.shape[1] != write.shape[0]:
raise ValueError(
f"Invalid SWD shapes: read={tuple(read.shape)}, write={tuple(write.shape)}"
)
self.read = nn.Parameter(read.detach().contiguous(), requires_grad=False)
self.write = nn.Parameter(write.detach().contiguous(), requires_grad=False)
self.bias = (
None
if bias is None
else nn.Parameter(bias.detach().contiguous(), requires_grad=False)
)
@property
def in_features(self) -> int:
return int(self.read.shape[0])
@property
def rank(self) -> int:
return int(self.read.shape[1])
@property
def out_features(self) -> int:
return int(self.write.shape[1])
def component_activations(self, inputs: torch.Tensor) -> torch.Tensor:
return inputs.matmul(self.read)
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
outputs = self.component_activations(inputs).matmul(self.write)
if self.bias is not None:
outputs = outputs + self.bias
return outputs
def _get_child(module: Any, name: str) -> Any:
if name.isdigit() and isinstance(module, (nn.ModuleList, nn.Sequential)):
return module[int(name)]
return getattr(module, name)
def _set_child(module: Any, name: str, value: nn.Module) -> None:
if name.isdigit() and isinstance(module, (nn.ModuleList, nn.Sequential)):
module[int(name)] = value
return
setattr(module, name, value)
def _resolve_parent(model: nn.Module, module_path: str) -> tuple[Any, str, nn.Module]:
parts = module_path.split(".")
if not parts or any(not part for part in parts):
raise ValueError(f"Invalid module path: {module_path!r}")
parent: Any = model
for part in parts[:-1]:
parent = _get_child(parent, part)
leaf = parts[-1]
target = _get_child(parent, leaf)
if not isinstance(target, nn.Module):
raise TypeError(f"Target at {module_path!r} is not an nn.Module")
return parent, leaf, target
def _module_device_dtype(module: nn.Module) -> tuple[torch.device, torch.dtype]:
weight = getattr(module, "weight", None)
if not torch.is_tensor(weight):
raise TypeError("Target module must expose a materialized weight tensor")
if weight.device.type == "meta":
raise ValueError("Load/materialize the base model before applying SWD factors")
if not weight.dtype.is_floating_point:
raise TypeError(f"Unsupported target weight dtype: {weight.dtype}")
return weight.device, weight.dtype
def _validate_base_weight(
module: nn.Module,
*,
module_path: str,
input_dim: int,
output_dim: int,
layout: str,
) -> None:
weight = getattr(module, "weight", None)
if not torch.is_tensor(weight):
raise TypeError(f"{module_path} does not expose a weight tensor")
expected = (
(input_dim, output_dim) if layout == "in_out" else (output_dim, input_dim)
)
if tuple(weight.shape) != expected:
raise ValueError(
f"Base weight mismatch at {module_path}: got {tuple(weight.shape)}, expected {expected}"
)
def _select_bias(
module: nn.Module,
tensors: dict[str, torch.Tensor],
spec: dict[str, Any],
*,
device: torch.device,
dtype: torch.dtype,
) -> torch.Tensor | None:
policy = spec["bias_policy"]
if policy == "checkpoint":
bias = tensors[spec["bias_key"]]
elif policy == "preserve_base":
bias = getattr(module, "bias", None)
elif policy == "none":
bias = None
else:
raise ValueError(f"Unknown bias policy: {policy}")
return None if bias is None else bias.detach().to(device=device, dtype=dtype)
def _fold_into_module(
module: nn.Module,
read: torch.Tensor,
write: torch.Tensor,
bias: torch.Tensor | None,
*,
layout: str,
) -> None:
dense_in_out = read.matmul(write)
dense = dense_in_out if layout == "in_out" else dense_in_out.transpose(0, 1)
weight = getattr(module, "weight")
with torch.no_grad():
weight.copy_(dense.to(device=weight.device, dtype=weight.dtype))
existing_bias = getattr(module, "bias", None)
if bias is not None:
if existing_bias is None:
raise ValueError("Checkpoint provides bias but base module has no bias parameter")
existing_bias.copy_(bias.to(existing_bias.device, existing_bias.dtype))
def load_swd_config(checkpoint_dir: str | Path) -> dict[str, Any]:
path = Path(checkpoint_dir) / "config.json"
with path.open(encoding="utf-8") as handle:
config = json.load(handle)
if config.get("schema_version") != "swd_factor_checkpoint_v1":
raise ValueError(f"Unsupported SWD schema: {config.get('schema_version')!r}")
return config
def apply_swd_checkpoint(
model: nn.Module,
checkpoint_dir: str | Path,
*,
mode: Literal["factorized", "folded"] = "factorized",
) -> list[AppliedReplacement]:
"""Apply one release checkpoint to a loaded base model.
``factorized`` installs :class:`SWDLinear` modules and preserves explicit
bottleneck activations. ``folded`` writes ``read @ write`` into the existing
dense modules for conventional inference.
"""
if mode not in {"factorized", "folded"}:
raise ValueError(f"Unknown mode: {mode}")
root = Path(checkpoint_dir)
config = load_swd_config(root)
tensors = load_file(root / config["weights_file"], device="cpu")
applied: list[AppliedReplacement] = []
for spec in config["module_replacements"]:
path = spec["module_path"]
parent, leaf, module = _resolve_parent(model, path)
input_dim = int(spec["input_dim"])
rank = int(spec["rank"])
output_dim = int(spec["output_dim"])
layout = spec["base_weight_layout"]
_validate_base_weight(
module,
module_path=path,
input_dim=input_dim,
output_dim=output_dim,
layout=layout,
)
device, dtype = _module_device_dtype(module)
read = tensors[spec["read_key"]]
write = tensors[spec["write_key"]]
if tuple(read.shape) != (input_dim, rank):
raise ValueError(f"Read tensor mismatch for {path}: {tuple(read.shape)}")
if tuple(write.shape) != (rank, output_dim):
raise ValueError(f"Write tensor mismatch for {path}: {tuple(write.shape)}")
read = read.to(device=device, dtype=dtype)
write = write.to(device=device, dtype=dtype)
bias = _select_bias(module, tensors, spec, device=device, dtype=dtype)
if mode == "factorized":
_set_child(parent, leaf, SWDLinear(read, write, bias))
else:
_fold_into_module(module, read, write, bias, layout=layout)
applied.append(
AppliedReplacement(path, input_dim, rank, output_dim, mode)
)
return applied