mergekit_config_gen.py

script

← Back to skill

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()