This repository contains the DRIFT-MEDIAN model-merging experiments built on top of FusionBench. DRIFT-MEDIAN merges task-specific models by resolving task-vector sign conflicts, keeping the largest task-vector updates per coordinate, and aggregating the remaining updates with a median-style rule.
In most experiments, DRIFT-MEDIAN is run through our modified version of Fisher
merging. The method still uses Fisher-style per-parameter importance weights, but the
final merge differs from vanilla Fisher averaging: task vectors are sparsified with
keep_ratio, optionally sign-resolved with elect_sign, and then combined with
median_or_mean=median or another supported aggregation mode.
fusion_bench/method/fisher_merging/: DRIFT-MEDIAN implementations.clip_fisher_merging.py: CLIP DRIFT-MEDIAN / modified Fisher merging.llm_fisher_merging.py: causal-LM DRIFT-MEDIAN / modified Fisher merging.gpt2_drift_median.py: GPT-2 DRIFT-MEDIAN.
config/method/fisher_merging/: Hydra method configs.*.shand*.py: validation, evaluation, ablation, and result processing scripts used for the paper experiments.docs/: upstream FusionBench documentation.
Run a single CLIP merge/evaluation with the modified Fisher merging implementation:
fusion_bench \
method=fisher_merging/clip_fisher_merging \
modelpool=CLIPVisionModelPool/clip-vit-base-patch32_TA8 \
taskpool=CLIPVisionModelTaskPool/clip-vit-classification_TA8 \
method.keep_ratio=1.0 \
method.scaling_factor=1.2 \
method.elect_sign=true \
method.median_or_mean=median \
method.use_fisher=true \
method.weighting_method=fisher \
report_save_path=outputs/clip_drift_median/report.jsonImportant knobs:
method.keep_ratio: fraction of task vectors kept per coordinate before merging.method.scaling_factor: multiplier for the merged task vector.method.elect_sign: enables sign election before the final merge.method.median_or_mean: one ofmean,medianmethod.use_fisher: whentrue, use Fisher weights; whenfalse, run the Fisher-free ablation.method.weighting_method: usuallyfisher; ablations includeabs_gradandabs_task_vector.
For the standard CLIP validation sweep, use:
bash validate_clip.shFor the corresponding test evaluation, use:
bash evaluate_clip.shFor causal language models, use method=fisher_merging/llm_fisher_merging. The main
experiments use the same DRIFT-MEDIAN idea through the modified Fisher merging path:
fusion_bench \
method=fisher_merging/llm_fisher_merging \
modelpool=CausalLMPool/mergebench/Llama-3.2-3B.yaml \
method.fisher_cache_path=/workspace/fisher_cache \
method.keep_ratio=0.8 \
method.scaling_factor=1.3 \
method.median_or_mean=median \
method.normalize_fisher_weight=false \
method.copy_base_embed_for_fisher=true \
method.merge_embed_layers=false \
report_save_path=outputs/llm_drift_median/report.json \
merged_model_save_path=outputs/llm_drift_median/modelThe paper-style wrappers are:
bash validate_llama_3b.sh
bash evaluate_llama_3b.sh
bash validate_llama_8b.sh
bash evaluate_llama_8b.shThese scripts assume the model checkpoints, Fisher caches, and output roots used in the experiment environment. Override paths with environment variables or edit the script headers before running on another machine.
GPT-2 uses its own DRIFT-MEDIAN class:
fusion_bench \
method=fisher_merging/gpt2_fisher_merging \
modelpool=gpt-2_glue \
taskpool=gpt-2_glue \
method.top_k=3 \
method.merge_lambda=1.35 \
method.num_fisher_examples=512 \
report_save_path=outputs/gpt2_drift_median/report.jsonValidation and evaluation helpers:
bash validate_gpt2.sh
bash gpt2_evaluate.shablation_validate_clip.sh: validates CLIP ablations.ablation_clip.sh: evaluates selected CLIP ablations.eval_wo_fisher_3b.sh: Fisher-free LLM evaluation.domain_change_clip.shandevaluate_domain_change_clip.sh: CLIP domain-change experiments.eval_calc_prr_*.py: post-processes reports and computes PRR summaries.find_best_valid_config_*.py: selects best validation configurations.
The mainline DRIFT-MEDIAN setting is the modified Fisher merge with
use_fisher=true, weighting_method=fisher, sign election enabled, and median
aggregation. Fisher-free and non-Fisher weighting modes are ablations.