Content hash: 339c3b92db7d0d5b011dcd95ab9694063522188442b6e6189c3a8ee57c2c07dc
#!/usr/bin/env python3
"""Generate mergekit YAML configs for common merge methods.
Supports: slerp (2 models), TIES, DARE-TIES, and task arithmetic.
Outputs a YAML file ready for `mergekit-yaml config.yml ./output`.
"""
from __future__ import annotations
import sys
from typing import Optional
def _model_sources(models):
return "\n".join(f" - model: {m}" for m in models)
MERGEKIT_TEMPLATES = {
"slerp": lambda base, models, params: f"""# SLERP merge: two models interpolated on the sphere
slices:
- sources:
- model: {models[0]}
layer_range: [0, {params['num_layers']}]
- model: {models[1]}
layer_range: [0, {params['num_layers']}]
merge_method: slerp
base_model: {base}
parameters:
t:
- filter: self_attn
value: [{params['t']}] * {params['num_layers']}
- filter: mlp
value: [{params['t']}] * {params['num_layers']}
tokenizer: {base}
dtype: bfloat16
""",
"ties": lambda base, models, params: f"""# TIES-Merging: Trim, Elect, Merge
slices:
- sources:
{_model_sources(models)}
merge_method: ties
base_model: {base}
parameters:
density: {params.get('density', 0.7)}
weight: {params.get('weight', 1.0)}
tokenizer: {base}
dtype: bfloat16
""",
"dare_ties": lambda base, models, params: f"""# DARE-TIES: Drop + TIES combine
slices:
- sources:
{_model_sources(models)}
merge_method: dare_ties
base_model: {base}
parameters:
density: {params.get('density', 0.7)}
weight: {params.get('weight', 1.0)}
int8_mask: {str(params.get('int8_mask', False)).lower()}
tokenizer: {base}
dtype: bfloat16
""",
"task_arithmetic": lambda base, models, params: f"""# Task Arithmetic: base + lambda * sum(deltas)
slices:
- sources:
{_model_sources(models)}
merge_method: task_arithmetic
base_model: {base}
parameters:
normalize: {str(params.get('normalize', False)).lower()}
weight: {params.get('weight', 0.3)}
tokenizer: {base}
dtype: bfloat16
""",
}
def generate_config(
method: str,
base: str,
models: list[str],
num_layers: int = 32,
**kwargs,
) -> str:
if method not in MERGEKIT_TEMPLATES:
valid = ", ".join(MERGEKIT_TEMPLATES)
raise ValueError(f"Unknown method '{method}'. Valid: {valid}")
params = dict(kwargs)
if "num_layers" not in params:
params["num_layers"] = num_layers
if "t" not in params:
params["t"] = 0.5
return MERGEKIT_TEMPLATES[method](base, models, params)
def main() -> None:
if len(sys.argv) < 3:
print("Usage: python mergekit_config_gen.py <method> <base_model> <model1> [model2...]")
print()
print("Methods: slerp (2 models only), ties, dare_ties, task_arithmetic")
print()
print("Examples:")
print(" python mergekit_config_gen.py slerp org/base org/code-7b org/math-7b")
print(" python mergekit_config_gen.py ties org/base org/code org/math org/reason")
sys.exit(2)
method = sys.argv[1]
base = sys.argv[2]
models = sys.argv[3:]
if method == "slerp" and len(models) != 2:
print("Error: slerp requires exactly 2 models", file=sys.stderr)
sys.exit(2)
config = generate_config(method=method, base=base, models=models)
output_file = f"mergecfg_{method}.yml"
with open(output_file, "w") as f:
f.write(config)
print(f"Wrote {output_file}")
print()
print("Run: mergekit-yaml", output_file, "./merged-output")
print("After: huggingface-cli upload myorg/merged-model ./merged-output")
if __name__ == "__main__":
main()