""" Model merging strategies for biological language models. Implements linear interpolation, SLERP, and task-vector merging. """ import os import torch import numpy as np import logging from typing import List, Dict, Optional from transformers import AutoModelForMaskedLM, AutoTokenizer, AutoConfig from safetensors.torch import save_file, load_file from pathlib import Path logger = logging.getLogger(__name__) class BioModelMerger: """Base class for biological model merging.""" def merge( self, parent_paths: List[str], save_path: str, base_model_path: Optional[str] = None, ) -> str: """Merge models and save result. Returns path to merged model.""" raise NotImplementedError class LinearMerge(BioModelMerger): """Simple linear interpolation of model weights.""" def __init__(self, weights: Optional[List[float]] = None): self.weights = weights # If None, uses equal weights def merge( self, parent_paths: List[str], save_path: str, base_model_path: Optional[str] = None, ) -> str: logger.info(f"Linear merging {len(parent_paths)} models into {save_path}") # Load state dicts states = [] for path in parent_paths: try: state = load_file(os.path.join(path, "model.safetensors")) except: # Fallback to PyTorch bin state = torch.load( os.path.join(path, "pytorch_model.bin"), map_location="cpu", ) states.append(state) # Determine weights n = len(states) if self.weights is None: weights = [1.0 / n] * n else: weights = self.weights[:n] # Normalize total = sum(weights) weights = [w / total for w in weights] # Merge merged = {} for key in states[0].keys(): merged[key] = sum(w * s[key] for w, s in zip(weights, states)) # Save os.makedirs(save_path, exist_ok=True) save_file(merged, os.path.join(save_path, "model.safetensors")) # Copy config and tokenizer from first parent self._copy_auxiliary_files(parent_paths[0], save_path) logger.info(f"Merged model saved to {save_path}") return save_path def _copy_auxiliary_files(self, source_path: str, dest_path: str) -> None: """Copy config, tokenizer files from source to destination.""" import shutil files_to_copy = [ "config.json", "tokenizer.json", "tokenizer_config.json", "vocab.txt", "special_tokens_map.json", "added_tokens.json", ] for fname in files_to_copy: src = os.path.join(source_path, fname) if os.path.exists(src): shutil.copy2(src, os.path.join(dest_path, fname)) class SlerpMerge(BioModelMerger): """Spherical linear interpolation for model weights.""" def __init__(self, t: float = 0.5): self.t = t def _slerp(self, v0: torch.Tensor, v1: torch.Tensor, t: float) -> torch.Tensor: """Spherical linear interpolation between two tensors.""" # Flatten for dot product v0_flat = v0.flatten().float() v1_flat = v1.flatten().float() # Normalize v0_norm = v0_flat / (v0_flat.norm() + 1e-8) v1_norm = v1_flat / (v1_flat.norm() + 1e-8) # Dot product dot = torch.clamp(torch.dot(v0_norm, v1_norm), -1.0, 1.0) # Angle theta = torch.acos(dot) if theta < 1e-6: # Linear fallback for very close vectors return (1 - t) * v0 + t * v1 # SLERP formula sin_theta = torch.sin(theta) w0 = torch.sin((1 - t) * theta) / sin_theta w1 = torch.sin(t * theta) / sin_theta result = w0 * v0 + w1 * v1 return result.to(v0.dtype) def merge( self, parent_paths: List[str], save_path: str, base_model_path: Optional[str] = None, ) -> str: if len(parent_paths) != 2: logger.warning("SLERP requires exactly 2 parents, falling back to linear merge") return LinearMerge().merge(parent_paths, save_path, base_model_path) logger.info(f"SLERP merging 2 models with t={self.t} into {save_path}") # Load state dicts states = [] for path in parent_paths: try: state = load_file(os.path.join(path, "model.safetensors")) except: state = torch.load( os.path.join(path, "pytorch_model.bin"), map_location="cpu", ) states.append(state) # SLERP merge merged = {} for key in states[0].keys(): merged[key] = self._slerp(states[0][key], states[1][key], self.t) # Save os.makedirs(save_path, exist_ok=True) save_file(merged, os.path.join(save_path, "model.safetensors")) self._copy_auxiliary_files(parent_paths[0], save_path) logger.info(f"SLERP merged model saved to {save_path}") return save_path def _copy_auxiliary_files(self, source_path: str, dest_path: str) -> None: import shutil files_to_copy = [ "config.json", "tokenizer.json", "tokenizer_config.json", "vocab.txt", "special_tokens_map.json", ] for fname in files_to_copy: src = os.path.join(source_path, fname) if os.path.exists(src): shutil.copy2(src, os.path.join(dest_path, fname)) class TaskVectorMerge(BioModelMerger): """Task vector arithmetic for model merging.""" def __init__(self, scaling: float = 0.5): self.scaling = scaling def merge( self, parent_paths: List[str], save_path: str, base_model_path: Optional[str] = None, ) -> str: """ Merge using task vector arithmetic: base + sum(scaling * (parent_i - base)) """ if base_model_path is None: # Use first parent as base base_model_path = parent_paths[0] logger.info(f"Task-vector merging with base={base_model_path}") # Load base try: base_state = load_file(os.path.join(base_model_path, "model.safetensors")) except: base_state = torch.load( os.path.join(base_model_path, "pytorch_model.bin"), map_location="cpu", ) # Compute task vectors and merge merged = {k: v.clone() for k, v in base_state.items()} for path in parent_paths: if path == base_model_path: continue try: state = load_file(os.path.join(path, "model.safetensors")) except: state = torch.load( os.path.join(path, "pytorch_model.bin"), map_location="cpu", ) for key in base_state.keys(): if key in state: task_vector = state[key] - base_state[key] merged[key] += self.scaling * task_vector # Save os.makedirs(save_path, exist_ok=True) save_file(merged, os.path.join(save_path, "model.safetensors")) # Copy config from base self._copy_auxiliary_files(base_model_path, save_path) logger.info(f"Task-vector merged model saved to {save_path}") return save_path def _copy_auxiliary_files(self, source_path: str, dest_path: str) -> None: import shutil files_to_copy = [ "config.json", "tokenizer.json", "tokenizer_config.json", "vocab.txt", ] for fname in files_to_copy: src = os.path.join(source_path, fname) if os.path.exists(src): shutil.copy2(src, os.path.join(dest_path, fname))