Optimization
Model Merging and Fusion ā Combining Knowledge Without Training
What if you could combine the strengths of multiple specialized models into one generalist model without additional training? Model merging achieves this by averaging, interpolating, or strategically combining model weights.
- Model Soups ā Averaging weights from fine-tuned models
- TIES-Merging ā Resolving interference through trim, elect, and disjoint merge
- DARE ā Drop and rescale for massive model merging
- Task Arithmetic ā Treating fine-tuning as task vectors
The whole can be greater than the sum of its parts if you know how to combine them.
Model Merging and Fusion
When you fine-tune a base model on different tasks, each specialized model learns task-specific knowledge in its weights. Model merging combines these specialized models into a single model that inherits the capabilities of all of them without requiring any training data or compute.
Model Soups
Linear Interpolation
import torch
import copy
from typing import Dict, List
class ModelSoupMerger:
"""Merge models by weight averaging."""
def __init__(self, models: List[Dict], weights: List[float] = None):
self.models = models
self.n_models = len(models)
if weights is None:
self.weights = [1.0 / self.n_models] * self.n_models
else:
total = sum(weights)
self.weights = [w / total for w in weights]
def merge(self) -> Dict:
"""Perform linear weight averaging."""
merged = {key: torch.zeros_like(v, dtype=torch.float32)
for key, v in self.models[0].items()}
for model, weight in zip(self.models, self.weights):
for key in merged:
merged[key] += weight * model[key].float()
return merged
Types of Model Soups
| Type | Description | When to Use |
|---|---|---|
| Uniform Soup | Equal weights for all models | Models are equally good |
| Greedy Soup | Add models only if they improve validation | Have validation data |
| Learned Soup | Optimize weights on validation set | Have enough data |
Task Arithmetic
Task Vectors
class TaskArithmetic:
"""Combine models using task vectors."""
def __init__(self, base_model, fine_tuned_models):
self.base_model = base_model
self.fine_tuned_models = fine_tuned_models
def compute_task_vectors(self):
"""Compute task vector for each fine-tuned model."""
task_vectors = []
for model in self.fine_tuned_models:
tv = {key: model[key] - self.base_model[key]
for key in self.base_model}
task_vectors.append(tv)
return task_vectors
def merge_with_task_arithmetic(self, scaling_factor=1.0):
"""Merge using task arithmetic."""
merged = copy.deepcopy(self.base_model)
for tv in self.compute_task_vectors():
for key in merged:
merged[key] += scaling_factor * tv[key]
return merged
def merge_with_negation(self, task_to_negate, scaling_factor=-1.0):
"""Remove a task's knowledge by negating its task vector."""
merged = copy.deepcopy(self.base_model)
task_vectors = self.compute_task_vectors()
for i, tv in enumerate(task_vectors):
factor = scaling_factor if i == task_to_negate else 1.0
for key in merged:
merged[key] += factor * tv[key]
return merged
TIES-Merging
The Interference Problem
TIES Algorithm
class TIESMerger:
"""TIES-Merging implementation."""
def __init__(self, base_model, fine_tuned_models, top_k=0.2):
self.base_model = base_model
self.fine_tuned_models = fine_tuned_models
self.top_k = top_k
def merge(self):
"""Perform TIES merging."""
task_vectors = self._compute_task_vectors()
trimmed = self._trim_updates(task_vectors)
consensus = self._elect_signs(trimmed)
return self._disjoint_merge(trimmed, consensus)
def _compute_task_vectors(self):
"""Compute task vectors for each model."""
return [{key: model[key] - self.base_model[key]
for key in self.base_model}
for model in self.fine_tuned_models]
def _trim_updates(self, task_vectors):
"""Keep only top-k% of updates by magnitude."""
trimmed = []
for tv in task_vectors:
trimmed_tv = {}
for key in tv:
flat = tv[key].flatten()
n_keep = int(len(flat) * self.top_k)
threshold = flat.abs().topk(n_keep).values[-1]
mask = tv[key].abs() >= threshold
trimmed_tv[key] = tv[key] * mask.float()
trimmed.append(trimmed_tv)
return trimmed
def _elect_signs(self, trimmed_vectors):
"""Determine consensus sign for each parameter."""
consensus = {}
for key in trimmed_vectors[0]:
sum_tv = torch.zeros_like(trimmed_vectors[0][key])
for tv in trimmed_vectors:
sum_tv += tv[key].sign()
consensus[key] = sum_tv.sign()
return consensus
def _disjoint_merge(self, trimmed_vectors, consensus_signs):
"""Merge only parameters where all models agree on sign."""
merged = copy.deepcopy(self.base_model)
for key in merged:
agreement = torch.ones_like(consensus_signs[key])
for tv in trimmed_vectors:
mask = (tv[key] != 0) & (tv[key].sign() == consensus_signs[key])
agreement *= mask.float()
sum_values = torch.zeros_like(merged[key])
count = torch.zeros_like(merged[key])
for tv in trimmed_vectors:
sum_values += tv[key] * agreement
count += (tv[key] != 0).float() * agreement
count = count.clamp(min=1)
merged[key] += sum_values / count
return merged
DARE (Drop And REscale)
Theory
class DAREMerger:
"""DARE (Drop and REscale) merging."""
def __init__(self, base_model, fine_tuned_models, drop_rate=0.9):
self.base_model = base_model
self.fine_tuned_models = fine_tuned_models
self.keep_prob = 1 - drop_rate
def merge(self, n_samples=10):
"""Merge with DARE (multiple samples for stability)."""
merged_models = [self._single_merge() for _ in range(n_samples)]
final = copy.deepcopy(merged_models[0])
for key in final:
for m in merged_models[1:]:
final[key] += m[key]
final[key] /= n_samples
return final
def _single_merge(self):
"""Single DARE merge with random dropping."""
merged = copy.deepcopy(self.base_model)
for model in self.fine_tuned_models:
for key in merged:
task_vec = model[key] - self.base_model[key]
mask = torch.bernoulli(
torch.full_like(task_vec, self.keep_prob)
)
dropped = task_vec * mask / self.keep_prob
merged[key] += dropped
return merged
SLERP (Spherical Linear Interpolation)
Comparison of Methods
| Method | Complexity | Interference Handling | Quality |
|---|---|---|---|
| Uniform Soup | O(n) | None | Good |
| Task Arithmetic | O(n) | Scaling factor | Good |
| TIES-Merging | O(nĀ·p) | Sign consensus | Excellent |
| DARE | O(nĀ·p) | Random dropping | Excellent |
| SLERP | O(n) | Pairwise only | Good |
Practical Guidelines
When to Use Each Method
| Scenario | Recommended Method |
|---|---|
| Models from same pre-trained | Uniform Soup |
| Different tasks, same base | Task Arithmetic |
| Many models (10+) | DARE |
| Conflicting tasks | TIES-Merging |
| Two models only | SLERP |
| Need best quality | TIES + DARE |
Example Workflow
def merge_pipeline(base_model, specialized_models, val_data=None):
"""Complete model merging workflow."""
# Step 1: Try uniform soup first
merger = ModelSoupMerger(specialized_models)
uniform_merge = merger.merge()
# Step 2: If validation data available, try greedy soup
if val_data:
greedy_merge = merger.merge_with_optimization(val_data)
# Step 3: For many models, try DARE
if len(specialized_models) > 5:
dare_merger = DAREMerger(base_model, specialized_models, drop_rate=0.9)
dare_merge = dare_merger.merge(n_samples=10)
# Step 4: For conflicting tasks, try TIES
ties_merger = TIESMerger(base_model, specialized_models, top_k=0.2)
ties_merge = ties_merger.merge()
return ties_merge # Usually best quality
Practice Exercises
-
Conceptual: Explain why fine-tuned models from the same pre-trained initialization can be averaged successfully. What assumption about the loss landscape makes this possible?
-
Mathematical: For 5 models each fine-tuned on different tasks, compute the number of parameters that need to be stored for TIES-Merging vs DARE merging.
-
Practical: Implement model soups by averaging weights from 3 LoRA-fine-tuned models and measure the performance on all three tasks.
-
Research: Compare TIES-Merging and DARE on merging 10 task-specific models. Which method better handles conflicting task requirements?
What to Learn Next
-> LoRA and PEFT Efficient fine-tuning using low-rank adaptation.
-> Knowledge Distillation for LLMs Training smaller models from larger teachers.
-> Low-Rank Factorization SVD decomposition and weight sharing techniques.
-> Quantization Techniques Deep Dive GPTQ, AWQ, GGUF, and INT4/INT8 methods.
-> Fine-Tuning LLMs Customizing language models for specific tasks.
-> Mixture of Experts Sparse architectures that scale efficiently.