splicing-bits / merge_engine.py
BluStatic's picture
Add merge_engine.py
4eedf45 verified
Raw
History Blame Contribute Delete
18.6 kB
"""
merge_engine.py – CPU-only shard-by-shard merge engine for Safetensors models.
Supports:
- Linear merge : merged = base * base_weight + donor * donor_weight
- DARE merge : delta masking with density + seed, then scaled addition
Design principles:
- Never load both complete models into RAM simultaneously
- Process and write tensors shard by shard
- Delete temporary files on completion or failure
- Preserve base model config / non-tensor files
- Generate valid model.safetensors.index.json when output is sharded
- No GPU, no training, no inference, no generation
"""
from __future__ import annotations
import json
import os
import shutil
import threading
import time
from pathlib import Path
from typing import Any, Callable, Generator
import numpy as np
import torch
from safetensors import safe_open
from safetensors.torch import save_file
# ---------------------------------------------------------------------------
# Cancellation token
# ---------------------------------------------------------------------------
class CancelToken:
"""Thread-safe cancellation flag."""
def __init__(self):
self._cancelled = threading.Event()
def cancel(self):
self._cancelled.set()
def is_cancelled(self) -> bool:
return self._cancelled.is_set()
def reset(self):
self._cancelled.clear()
# ---------------------------------------------------------------------------
# Low-level tensor helpers (CPU only)
# ---------------------------------------------------------------------------
def _open_st(path: str | Path):
return safe_open(str(path), framework="pt", device="cpu")
def _linear_merge(
t_base: torch.Tensor,
t_donor: torch.Tensor,
base_weight: float,
donor_weight: float,
) -> torch.Tensor:
"""merged = base * base_weight + donor * donor_weight"""
dtype = t_base.dtype
result = t_base.float() * base_weight + t_donor.float() * donor_weight
return result.to(dtype)
def _dare_merge(
t_base: torch.Tensor,
t_donor: torch.Tensor,
base_weight: float,
donor_weight: float,
density: float,
rng: np.random.Generator,
) -> torch.Tensor:
"""
DARE merge:
1. delta = donor - base
2. mask = Bernoulli(density) over delta shape
3. delta_masked = delta * mask
4. delta_rescaled = delta_masked / density (expectation-preserving)
5. merged = base + donor_weight * delta_rescaled
(base_weight applied to base component implicitly via the formula)
When density == 1.0 this reduces to:
merged = base + donor_weight * (donor - base)
= base * (1 - donor_weight) + donor * donor_weight
which is equivalent to linear with base_weight = 1 - donor_weight.
"""
dtype = t_base.dtype
base_f = t_base.float()
donor_f = t_donor.float()
delta = donor_f - base_f
flat_mask = rng.binomial(1, density, size=delta.numel()).astype(np.float32)
mask = torch.from_numpy(flat_mask).reshape(delta.shape)
if density > 0.0:
delta_rescaled = (delta * mask) / density
else:
delta_rescaled = torch.zeros_like(delta)
merged = base_f * base_weight + donor_weight * delta_rescaled
return merged.to(dtype)
# ---------------------------------------------------------------------------
# Shard discovery helpers
# ---------------------------------------------------------------------------
NON_TENSOR_EXTENSIONS = {
".json", ".txt", ".md", ".yaml", ".yml",
".py", ".model", ".tiktoken", ".vocab",
".merges", ".special_tokens_map",
}
TENSOR_EXTENSIONS = {".safetensors"}
def _find_safetensors_shards(folder: Path) -> list[Path]:
"""Return sorted list of .safetensors files in folder, excluding index files."""
shards = sorted(
p for p in folder.glob("*.safetensors")
if "index" not in p.name.lower()
)
return shards
def _read_index(folder: Path) -> dict | None:
"""Read model.safetensors.index.json if present."""
idx_path = folder / "model.safetensors.index.json"
if idx_path.exists():
with open(idx_path) as f:
return json.load(f)
return None
def _build_key_to_shard_map(folder: Path) -> dict[str, str]:
"""
Return {tensor_key: shard_filename} by reading the index file
or by scanning all shards directly.
"""
idx = _read_index(folder)
if idx and "weight_map" in idx:
return idx["weight_map"]
# No index – scan shards
key_map: dict[str, str] = {}
for shard in _find_safetensors_shards(folder):
with _open_st(shard) as f:
for k in f.keys():
key_map[k] = shard.name
return key_map
def _copy_non_tensor_files(src: Path, dst: Path, log: Callable[[str], None]):
"""Copy config / tokenizer / non-weight files from src to dst."""
dst.mkdir(parents=True, exist_ok=True)
for item in src.iterdir():
if item.is_dir():
continue
if item.suffix in TENSOR_EXTENSIONS:
continue
if item.name.startswith("."):
continue
dest_file = dst / item.name
shutil.copy2(item, dest_file)
log(f" Copied config file: {item.name}")
# ---------------------------------------------------------------------------
# Main merge engine
# ---------------------------------------------------------------------------
MAX_SHARD_BYTES = 4 * 1024 ** 3 # 4 GB per output shard
class MergeEngine:
"""
Orchestrates a shard-by-shard CPU merge of two Safetensors model folders
(or single-file models).
Parameters
----------
base_folder : Path to the downloaded base model folder
donor_folder : Path to the downloaded donor model folder
output_folder : Path where merged output will be written
method : "linear" or "dare"
base_weight : Weight applied to base tensors
donor_weight : Weight applied to donor tensors / delta
dare_density : Fraction of delta elements to keep (DARE only)
seed : Random seed for reproducibility (DARE only)
log_fn : Callable(str) for progress messages
cancel_token : CancelToken instance
"""
def __init__(
self,
base_folder: Path,
donor_folder: Path,
output_folder: Path,
method: str,
base_weight: float,
donor_weight: float,
dare_density: float,
seed: int,
log_fn: Callable[[str], None],
cancel_token: CancelToken,
):
self.base_folder = Path(base_folder)
self.donor_folder = Path(donor_folder)
self.output_folder = Path(output_folder)
self.method = method.lower()
self.base_weight = base_weight
self.donor_weight = donor_weight
self.dare_density = dare_density
self.seed = seed
self.log = log_fn
self.cancel = cancel_token
self._rng: np.random.Generator | None = None
# ------------------------------------------------------------------
# Public entry point
# ------------------------------------------------------------------
def run(self) -> dict[str, Any]:
"""
Execute the merge. Returns a recipe dict describing what was done.
Raises RuntimeError on hard failure.
Cleans up temp files on failure.
"""
self.output_folder.mkdir(parents=True, exist_ok=True)
if self.method == "dare":
self._rng = np.random.default_rng(self.seed)
try:
recipe = self._merge()
except Exception:
self._cleanup_output()
raise
return recipe
# ------------------------------------------------------------------
# Internal merge logic
# ------------------------------------------------------------------
def _merge(self) -> dict[str, Any]:
self.log("Building tensor → shard maps …")
base_key_map = _build_key_to_shard_map(self.base_folder)
donor_key_map = _build_key_to_shard_map(self.donor_folder)
all_base_keys = set(base_key_map.keys())
all_donor_keys = set(donor_key_map.keys())
common_keys = all_base_keys & all_donor_keys
base_only_keys = all_base_keys - all_donor_keys
self.log(
f"Keys – base: {len(all_base_keys)}, "
f"donor: {len(all_donor_keys)}, "
f"common: {len(common_keys)}, "
f"base-only (pass-through): {len(base_only_keys)}"
)
if not common_keys:
raise RuntimeError(
"No common tensor keys found. Models are incompatible."
)
# Group base shards so we process one shard at a time
shard_to_keys: dict[str, list[str]] = {}
for k in all_base_keys:
shard = base_key_map[k]
shard_to_keys.setdefault(shard, []).append(k)
# We'll accumulate output tensors and flush when shard is full
output_weight_map: dict[str, str] = {} # key → output shard filename
output_shard_idx = 0
pending_tensors: dict[str, torch.Tensor] = {}
pending_bytes = 0
def flush_pending():
nonlocal output_shard_idx, pending_tensors, pending_bytes
if not pending_tensors:
return
shard_name = f"model-{output_shard_idx + 1:05d}-of-XXXXX.safetensors"
out_path = self.output_folder / shard_name
self.log(f" Writing output shard: {shard_name} ({len(pending_tensors)} tensors)")
save_file(pending_tensors, str(out_path), metadata={"format": "pt"})
for k in pending_tensors:
output_weight_map[k] = shard_name
output_shard_idx += 1
pending_tensors = {}
pending_bytes = 0
total_shards = len(shard_to_keys)
for shard_idx, (base_shard_name, shard_keys) in enumerate(
sorted(shard_to_keys.items()), start=1
):
if self.cancel.is_cancelled():
raise RuntimeError("Merge cancelled by user.")
self.log(
f"Processing base shard {shard_idx}/{total_shards}: "
f"{base_shard_name} ({len(shard_keys)} tensors)"
)
base_shard_path = self.base_folder / base_shard_name
# Collect which donor shards we need for this batch of keys
donor_shards_needed: dict[str, list[str]] = {}
for k in shard_keys:
if k in common_keys:
ds = donor_key_map[k]
donor_shards_needed.setdefault(ds, []).append(k)
# Load donor tensors we need (one donor shard at a time)
donor_tensors: dict[str, torch.Tensor] = {}
for donor_shard_name, needed_keys in donor_shards_needed.items():
if self.cancel.is_cancelled():
raise RuntimeError("Merge cancelled by user.")
donor_shard_path = self.donor_folder / donor_shard_name
self.log(f" Loading donor shard: {donor_shard_name}")
with _open_st(donor_shard_path) as df:
for k in needed_keys:
donor_tensors[k] = df.get_tensor(k)
# Now process base shard tensor by tensor
with _open_st(base_shard_path) as bf:
for k in sorted(shard_keys):
if self.cancel.is_cancelled():
raise RuntimeError("Merge cancelled by user.")
t_base = bf.get_tensor(k)
if k in common_keys:
t_donor = donor_tensors[k]
# Cast donor to base dtype
t_donor = t_donor.to(t_base.dtype)
if self.method == "linear":
merged = _linear_merge(
t_base, t_donor,
self.base_weight, self.donor_weight,
)
elif self.method == "dare":
merged = _dare_merge(
t_base, t_donor,
self.base_weight, self.donor_weight,
self.dare_density, self._rng,
)
else:
raise ValueError(f"Unknown merge method: {self.method}")
else:
# Base-only key: pass through unchanged
merged = t_base
pending_tensors[k] = merged
pending_bytes += merged.numel() * merged.element_size()
if pending_bytes >= MAX_SHARD_BYTES:
flush_pending()
# Free donor tensors for this batch
donor_tensors.clear()
# Flush remaining tensors
flush_pending()
# Fix up shard filenames now that we know the total count
total_out_shards = output_shard_idx
self.log(f"Renaming output shards (total: {total_out_shards}) …")
new_weight_map: dict[str, str] = {}
for old_name in sorted(set(output_weight_map.values())):
# Extract index from "model-00001-of-XXXXX.safetensors"
parts = old_name.split("-")
idx = int(parts[1])
new_name = f"model-{idx:05d}-of-{total_out_shards:05d}.safetensors"
old_path = self.output_folder / old_name
new_path = self.output_folder / new_name
old_path.rename(new_path)
for k, v in output_weight_map.items():
if v == old_name:
new_weight_map[k] = new_name
# Write index if more than one shard
if total_out_shards > 1:
self._write_index(new_weight_map)
elif total_out_shards == 1:
# Rename single shard to canonical name
single_old = self.output_folder / list(new_weight_map.values())[0]
single_new = self.output_folder / "model.safetensors"
single_old.rename(single_new)
new_weight_map = {k: "model.safetensors" for k in new_weight_map}
# Copy non-tensor files from base
self.log("Copying config / tokenizer files from base model …")
_copy_non_tensor_files(self.base_folder, self.output_folder, self.log)
recipe = {
"method": self.method,
"base_weight": self.base_weight,
"donor_weight": self.donor_weight,
"dare_density": self.dare_density if self.method == "dare" else None,
"seed": self.seed if self.method == "dare" else None,
"common_keys": len(common_keys),
"base_only_keys": len(base_only_keys),
"output_shards": total_out_shards,
}
return recipe
# ------------------------------------------------------------------
# Index writing
# ------------------------------------------------------------------
def _write_index(self, weight_map: dict[str, str]):
idx = {
"metadata": {"format": "pt"},
"weight_map": weight_map,
}
idx_path = self.output_folder / "model.safetensors.index.json"
with open(idx_path, "w") as f:
json.dump(idx, f, indent=2)
self.log(f"Wrote index: {idx_path.name}")
# ------------------------------------------------------------------
# Cleanup
# ------------------------------------------------------------------
def _cleanup_output(self):
"""Remove partial output on failure."""
if self.output_folder.exists():
shutil.rmtree(self.output_folder, ignore_errors=True)
self.log("Cleaned up partial output folder.")
# ---------------------------------------------------------------------------
# Single-file merge (LoRA or small model in one .safetensors file)
# ---------------------------------------------------------------------------
def merge_single_files(
base_path: Path,
donor_path: Path,
output_path: Path,
method: str,
base_weight: float,
donor_weight: float,
dare_density: float,
seed: int,
log_fn: Callable[[str], None],
cancel_token: CancelToken,
) -> dict[str, Any]:
"""
Merge two single .safetensors files (e.g. LoRAs) tensor by tensor.
Writes a single output .safetensors file.
"""
rng = np.random.default_rng(seed) if method == "dare" else None
with _open_st(base_path) as bf, _open_st(donor_path) as df:
keys_base = set(bf.keys())
keys_donor = set(df.keys())
common = keys_base & keys_donor
base_only = keys_base - keys_donor
if not common:
raise RuntimeError("No common tensor keys. Files are incompatible.")
log_fn(f"Merging {len(common)} common tensors, {len(base_only)} pass-through …")
merged: dict[str, torch.Tensor] = {}
total = len(keys_base)
done = 0
for k in sorted(keys_base):
if cancel_token.is_cancelled():
raise RuntimeError("Merge cancelled by user.")
t_base = bf.get_tensor(k)
if k in common:
t_donor = df.get_tensor(k).to(t_base.dtype)
if method == "linear":
merged[k] = _linear_merge(t_base, t_donor, base_weight, donor_weight)
elif method == "dare":
merged[k] = _dare_merge(t_base, t_donor, base_weight, donor_weight,
dare_density, rng)
else:
raise ValueError(f"Unknown method: {method}")
else:
merged[k] = t_base
done += 1
if done % 50 == 0 or done == total:
log_fn(f" {done}/{total} tensors processed …")
output_path.parent.mkdir(parents=True, exist_ok=True)
log_fn(f"Saving merged file: {output_path.name}")
save_file(merged, str(output_path), metadata={"format": "pt"})
return {
"method": method,
"base_weight": base_weight,
"donor_weight": donor_weight,
"dare_density": dare_density if method == "dare" else None,
"seed": seed if method == "dare" else None,
"common_keys": len(common),
"base_only_keys": len(base_only),
"output_shards": 1,
}