Spaces:
Paused
Paused
| """ | |
| 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, | |
| } |