FlagEmbedding 的 LM-Cocktail 模型融合实战指南:用模型合并低成本提升大模型与 Embedding 模型性能
【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding
导读
LM-Cocktail 是 BAAI FlagEmbedding 开源项目中的一个模型融合(Model Merging)工具库,它把"像调制鸡尾酒一样"按比例混合多个同架构模型参数这一思路,落地为开箱即用的 Python API。本文以 research/LM_Cocktail/README.md 为核心,结合仓库内的 cocktail.py、utils.py 与示例数据文件,系统讲解 LM-Cocktail 的原理、三个核心函数的完整参数语义与调用示例、LLM/Embedding/Reranker 三类模型的使用差异,以及如何用官方脚本复现论文中的性能结果。读完本文,你将能够不经过任何额外训练,用几行代码缓解微调带来的灾难性遗忘、为全新任务"零训练"生成专用模型,甚至近似实现多任务学习。
LM-Cocktail 是什么:以"模型合并"代替"重新训练"
LM-Cocktail(论文标题Resilient Tuning of Language Models via Model Merging,arXiv:2311.13534)的核心思想非常朴素:把多个具有相同架构、相同初始化参数(same architecture and same initialization parameter)的模型,按照一组权重做参数级加权平均,得到一个新的模型。它把微调语言模型的过程类比为调制一杯鸡尾酒——基酒(Base Model)与各种风味利口酒(在不同任务上微调得到的模型)按比例混合,就能得到想要的"口感"。
从源码来看,这一思想在 cocktail.py 中体现为清晰的实现链路:
get_model_param_list逐个加载模型并收集state_dict();merge_param对每个参数张量执行new_param[k] = Σ w_i * param_i[k](整型张量如int64/int32直接取第一个模型的值,不参与平均,见 utils.py);- 将融合后的参数
load_state_dict回第一个模型并保存。
在 FlagEmbedding 生态中,LM-Cocktail 的典型应用场景是:
- 缓解灾难性遗忘(Catastrophic Forgetting):微调后模型在目标任务上变强,却在其他任务上大幅退化,通过混合微调模型与基座模型找回通用能力;
- 新任务零训练(New Task without Fine-tuning):用少量示例数据为不同微调模型自动计算权重并合并,直接生成针对新任务的模型;
- 近似多任务学习(Approximate Multitask Learning):合并多个单任务微调模型,得到一个能同时处理多个任务的模型。
使用前提:被合并的模型必须架构相同且初始化参数相同,否则参数级平均在数学上没有意义。
安装与包结构
LM-Cocktail 以独立 Python 包的形式发布,支持两种安装方式。
推荐从源码安装(获取最新版本):
git clone https://github.com/FlagOpen/FlagEmbedding.git cd FlagEmbedding/research/LM_Cocktail pip install -e .也可以直接用 pip 安装:
pip install -U LM_Cocktail包的核心代码位于 research/LM_Cocktail/LM_Cocktail 目录,结构如下:
__init__.py:导出全部三个公开函数mix_models、mix_models_with_data、mix_models_by_layers(见init.py);cocktail.py:三个顶层融合函数及 Encoder 模型的 sentence-transformers 格式转换辅助函数save_ckpt_for_sentence_transformers;utils.py:模型加载、参数收集、参数加权平均、基于示例数据计算权重等底层实现。
setup.py 中声明的运行依赖为torch>=1.6.0、transformers>=4.18.0、datasets与accelerate>=0.20.1。此外mix_models_by_layers还依赖 Hugging Faceaccelerate的init_empty_weights与load_checkpoint_and_dispatch实现低内存逐层加载。
核心函数一:mix_models——按指定权重合并模型
mix_models是 LM-Cocktail 最基本也最常用的函数:按照用户给定的权重列表,把多个模型做加权平均合并。典型用途就是微调后混合微调模型与基座模型,缓解灾难性遗忘。
函数签名与参数语义
mix_models(model_names_or_paths: List[str], model_type: str, weights: List[float], output_path: str = None)| 参数 | 类型 | 含义与约束 |
|---|---|---|
model_names_or_paths | List[str] | 待合并模型的名字(Hugging Face Hub)或本地路径列表,长度必须与weights一致 |
model_type | str | 模型类型,取值"decoder"/"encoder"/"reranker" |
weights | List[float] | 每个模型的融合权重,所有权重之和必须等于 1(源码中通过assert sum(weights) - 1 <= 1e-3校验) |
output_path | str | 合并后模型的保存路径;为None时仅返回内存中的模型,不落盘 |
三类模型的使用示例
合并 LLM(decoder 类型)——例如用 0.7 的权重保留基础能力、0.3 的权重引入目标任务专长,在通用性与专业性之间取折中:
from LM_Cocktail import mix_models, mix_models_with_data # mix LLMs and save it to output_path: ./mixed_model_1 model = mix_models( model_names_or_paths=["meta-llama/Llama-2-7b-chat-hf", "Shitao/llama2-ag-news"], model_type='decoder', weights=[0.7, 0.3], output_path='./mixed_llm') # you can select a weight for your models to get a trade-off between generality and expertise.合并 Embedding 模型(encoder 类型):
# Mix Embedding Models model = mix_models( model_names_or_paths=["BAAI/bge-base-en-v1.5", "Shitao/bge-hotpotqa"], model_type='encoder', weights=[0.5, 0.5], output_path='./mixed_embedder')合并 Reranker 模型(reranker 类型):
# Mix reranker Models model = mix_models( model_names_or_paths=["BAAI/bge-reranker-base", "BAAI/bge-reranker-base"], model_type='reranker', weights=[0.5, 0.5], output_path="./mixed_reranker")多模型合并
mix_models并不限制只合并两个模型,可以一次性合并任意多个:
model = mix_models( model_names_or_paths=["BAAI/bge-base-en-v1.5", "Shitao/bge-hotpotqa", "Shitao/bge-quora", "Shitao/bge-msmarco"], model_type='encoder', weights=[0.3, 0.2, 0.2, 0.3], output_path='./mixed_embedder_2') # The sum of weights should be equal to 1.源码行为细节(值得注意的几点)
- 整型参数不平均:在 utils.py 的
merge_param中,torch.int64/torch.int32类型的参数(如 embedding 层中某些整数索引类参数)直接沿用第一个模型的值,不做加权求和; - 权重回显:合并完成后会在终端打印每个模型及其对应权重("weight for each model"),便于核对;
- Encoder 特殊处理:当
model_type == "encoder"且指定了output_path时,保存完model.save_pretrained与tokenizer.save_pretrained之后,还会调用save_ckpt_for_sentence_transformers,把模型转换成sentence-transformers 格式(默认pooling_method='cls'、normalized=True),因此合并出的 Embedding 模型可以直接被SentenceTransformer加载使用(见 cocktail.py)。
核心函数二:mix_models_with_data——用少量示例自动计算权重并合并
mix_models_with_data解决的是"不知道每个模型该配多少权重"的问题:它在给定的小规模示例数据上分别计算每个模型的损失,再根据损失自动分配融合权重,然后完成合并。它既可以用于为新任务免训练生成模型,也可以用于融合其他任务上已有的微调模型来增强自己的模型。
函数签名与参数语义
mix_models_with_data(model_names_or_paths: List[str], model_type: str, example_data: List[Dict], temperature: float = 5.0, batch_size: int = 2, max_input_length: int = 2048, neg_number: int = 7, output_path: str = None)| 参数 | 类型 | 默认值 | 含义 |
|---|---|---|---|
model_names_or_paths | List[str] | - | 候选模型列表 |
model_type | str | - | 取值"decoder"/"encoder"/"encoder-decoder" |
example_data | List[Dict] | - | 示例数据,格式随模型类型而变(见下文) |
temperature | float | 5.0 | 温度参数,用于调节融合权重的分布:温度越高,权重分布越均匀;越低则越"尖锐",对低损失模型更倾斜 |
batch_size | int | 2 | 计算损失时的批大小 |
max_input_length | int | 2048 | 输入的最大 token 数 |
neg_number | int | 7 | 计算 Embedding 对比损失时每个 query 使用的负样本个数 |
output_path | str | None | 保存路径 |
示例数据格式
LLM(decoder)的示例数据格式:一个字典列表,每个字典形如:
{"input": str, "output": str}LM-Cocktail 会计算 output 部分的语言模型损失(源码中llm_loss对 prompt 部分 token 置为-100,只对输出部分计算损失,见 utils.py):
example_data = [ {"input": "Question: when was the last time anyone was on the moon? Answer:\n", "output": "14 December 1972 UTC"}, {"input": "Review: \"it 's a charming and often affecting journey . \" Is this movie review sentence negative or positive?\n", "output": "Positive"} ] model = mix_models_with_data( model_names_or_paths=["meta-llama/Llama-2-7b-chat-hf", "Shitao/llama2-ag-news", "Shitao/llama2-nq"], model_type='decoder', example_data=example_data, temperature=5.0) # you can set the temperature argument to adjust the distribution of mixing weightsEmbedding(encoder)的示例数据格式:每个字典形如:
{"query": str, "pos": List[str], 'neg': List[str]}其中pos是正例文本列表,neg是负例文本列表。LM-Cocktail 会为每个 query 从负例中随机抽取neg_number个,与正例一起构造对比学习损失(embedder_loss中把 query 与 passage 的 embedding 做点积后除以温度 0.05,再计算交叉熵,见 utils.py):
example_data = [ {"query": "How does one become an actor in the Telugu Film Industry?", "pos": [" How do I become an actor in Telugu film industry?"], "neg": [" What is the story of Moses and Ramesses?", " Does caste system affect economic growth of India?"]}, {"query": "Why do some computer programmers develop amazing software or new concepts, while some are stuck with basic programming work?", "pos": [" Why do some computer programmers develops amazing softwares or new concepts, while some are stuck with basics programming works?"], "neg": [" When visiting a friend, do you ever think about what would happen if you did something wildly inappropriate like punch them or destroy their furniture?", " What is the difference between a compliment and flirting?"]} ] model = mix_models_with_data( model_names_or_paths=["BAAI/bge-base-en-v1.5", "Shitao/bge-hotpotqa", "Shitao/bge-quora"], model_type='encoder', example_data=example_data, temperature=5.0, max_input_length=512, neg_number=2)权重计算原理
从 utils.py 的compute_weights可以还原权重计算过程:
- 将所有候选模型的参数依次
load_state_dict到基座模型上,在torch.no_grad()下用同一批示例数据计算各自的损失example_loss; - 按公式
weights = softmax(-loss / temperature)得到最终权重——损失越小的模型(在示例数据上表现越好)获得越大的权重; - 计算在 CUDA / NPU / CPU 上自动选择设备(
torch.cuda.is_available()优先,其次is_torch_npu_available())。
核心函数三:mix_models_by_layers——逐层合并以降低内存占用
对于 7B 乃至更大规模的模型,mix_models需要一次性把所有模型的完整state_dict载入内存,代价很高。mix_models_by_layers通过创建临时目录存放各模型的权重,再逐层加载、逐层融合来降低峰值内存:
from LM_Cocktail import mix_models_by_layers # Mix Large Language Models (LLMs) and save the combined model to the path: ./mixed_llm model = mix_models_by_layers( model_names_or_paths=["meta-llama/Llama-2-7b-chat-hf", "Shitao/llama2-ag-news"], model_type='decoder', weights=[0.7, 0.3], output_path='./mixed_llm')其源码实现(cocktail.py)分四步:
get_model_param_dirs在临时目录中为每个模型建立子目录,把各层参数逐个torch.save为独立的.ckpt文件,随后把模型置为meta设备并gc.collect()释放内存(见 utils.py);merge_param_by_layer逐参数文件读取、加权求和,边加载边合并边释放,最后写回一个临时 checkpoint 文件(utils.py);- 用
accelerate.init_empty_weights创建 meta 设备上的空模型骨架,再通过load_checkpoint_and_dispatch从临时 checkpoint 逐层装载,并调用model.tie_weights()保证权重共享层一致; - 合并完成后自动删除临时 checkpoint 文件与临时目录。
注意:
mix_models_by_layers的model_type仅支持"decoder"/"encoder"/"reranker",不支持"encoder-decoder"(源码中对未知类型抛出NotImplementedError)。这与mix_models_with_data支持 seq2seq 模型略有不同。
应用场景详解:三种典型用法
场景一:缓解灾难性遗忘
微调基座 LLM 通常会导致目标领域之外通用能力的严重退化。官方推荐的做法分两种:
- 两模型混合:直接用
mix_models把微调模型与基座模型按比例混合(例如 0.7 : 0.3),即可在显著提升目标任务表现的同时维持其他无关任务的能力; - 多模型增强:如果社区或团队里已存在在其他任务上微调好的模型,先收集 5 条自己任务的示例数据,用
mix_models_with_data计算权重并合并这些可用模型——该函数会自动给低质量模型分配更低的权重,避免拖累目标任务;最后再用mix_models把产物与自己的微调模型合并。
场景二:新任务免训练增强
当面对一个从未微调过的全新任务时,不需要训练模型:只需给出少量示例数据(例如 5 条),配合一批来自开源社区或其他任务的可用模型,mix_models_with_data会根据示例损失自动为不同模型分配不同的融合权重,生成一个任务特化的新模型。
场景三:近似多任务学习
若手头已有多个在不同任务上微调过的模型,可以将它们合并为一个模型来近似多任务学习的效果,合并后的模型能够同时承担多个任务。这本质上是把"训练一个多任务模型"的成本,替换为"合并几个单任务模型"的廉价操作。
性能表现:官方实验数据
LM-Cocktail 论文报告了以下关键结果(来源:research/LM_Cocktail/README.md,详细数据请参考论文正文):
灾难性遗忘场景(LLM)
| 模型 | 目标任务 | 其他任务(29 个) |
|---|---|---|
| Llama | 40.8 | 46.8 |
| Fine-tuned | 94.4 | 38.6 |
| LM-Cocktail(2 模型)[1] | 94.5 | 47.7 |
| LM-Cocktail(10 模型)[2] | 94.4 | 48.3 |
[1]:合并 2 个模型,即微调模型与基座模型;[2]:基于 5 条示例合并 10 个模型,即微调模型、基座模型与其他任务上微调的 8 个模型。
可见,微调把通用能力从 46.8 拉低到 38.6,而 LM-Cocktail 在几乎不损失目标任务(94.4 → 94.5)的情况下,把其他任务分数恢复到 47.7~48.3,甚至略超原始基座。
灾难性遗忘场景(Embedding)
| 模型 | 目标任务 | 其他任务(14 个) |
|---|---|---|
| BGE | 71.8 | 49.8 |
| Fine-tuned | 76.0 | 48.5 |
| LM-Cocktail(2 模型) | 74.8 | 50.0 |
| LM-Cocktail(10 模型) | 74.7 | 50.6 |
新任务免训练场景
基于 5 条示例合并 10 个在其他任务上微调的模型:
| 模型 | MMLU(57 个任务) |
|---|---|
| Llama | 45.9 |
| Llama-5shot | 46.7 |
| LM-Cocktail(10 模型) | 48.0 |
| 模型 | Retrieval(12 个任务) |
|---|---|
| BGE | 47.3 |
| LM-Cocktail(10 模型) | 48.8 |
以上数字为论文复现实验中的报告结果,复现时请以论文与仓库脚本为准;不同数据集版本或合并模型组合下数值可能略有波动。
仓库中的实验说明图 research/LM_Cocktail/images/pic.png 以图表形式展示了上述思路:左侧为 Llama / Fine-tuned / LM-Cocktail 在目标任务与其他任务上的性能对比柱状图,右侧为"目标任务微调"与"LM-Cocktail 多模型融合"两条处理路径及模型融合架构示意。
复现与评估实验
复现 LLM 的实验结果
官方提供了可复现的数据与脚本:
- 模型:作者在 9 个任务上微调了
meta-llama/Llama-2-7b-chat-hf,微调模型可在Shitao账号下找到。注意其中大多数微调模型在其他无关任务上表现较差(这正是灾难性遗忘现象); - 示例数据:llm_examples.json(涵盖 mnli_m、mrpc、natural_questions、squad_v1、sst2、winogrande、ag_news、common_gen、hellaswag 等任务的示例,格式即上文所述的
{"input", "output"}结构); - 评估脚本:使用仓库内 llm_embedder 的评估文档 中描述的脚本,例如:
# for 30 tasks from FLAN torchrun --nproc_per_node 8 -m evaluation.eval_icl \ --retrieval_method no \ --few_shot 0 \ --data_root /data/llm-embedder \ --model_name_or_path ./mixed_model_1 # for MMLU datasets torchrun --nproc_per_node 8 -m evaluation.eval_mmlu \ --retrieval_method no \ --few_shot 0 \ --data_root /data/llm-embedder \ --model_name_or_path ./mixed_model_2(MMLU 数据集来自cais/mmlu,评估时使用 dev 集的示例做 in-context learning。)
复现 Embedding 模型的实验结果
- 模型:作者在 9 个任务上微调了
bge-base-en-v1.5,微调模型同样可在Shitao账号下找到; - 示例数据:embedder_examples.json(按任务分组,包含 ArguAna、ClimateFEVER、CQADupstack 系列、DBPedia、FEVER、HotpotQA、NFCorpus、NQ、QuoraRetrieval、SCIDOCS、SciFact、TRECCOVID、Touche2020 等检索任务的
{"query", "pos", "neg"}示例); - 评估脚本:使用 C_MTEB 目录 中的 MTEB 评估脚本:
python eval_MTEB.py --model_name_or_path mixed_model --task_type Retrieval实践要点与注意事项
- 同构同源是前提:合并的模型必须架构相同、初始化相同。README 明确指出 "the models used to merge need to have the same architecture and the same initialization parameter"。
- 权重和必须为 1:
mix_models与mix_models_by_layers均有assert sum(weights) - 1 <= 1e-3的强校验;mix_models_with_data的权重由损失自动归一化,无需手工保证。 - 温度参数控制权重分布:
temperature越大,softmax(-loss/temperature)得到的权重越接近均匀分布;温度越小,表现最好的模型权重越突出。在候选模型质量参差、示例数据可信时,可调低温度以获得更明确的"择优"效果。 - 示例数据贵精不贵多:官方实验表明几条(例如 5 条)精心挑选的示例数据即可完成权重计算,这正是"免训练"特性的关键。
- 内存敏感时使用逐层合并:大模型优先选择
mix_models_by_layers,它会通过临时目录 + 逐层装载显著降低峰值内存,并在结束后自动清理临时文件。 - Embedding 模型输出为 sentence-transformers 格式:
mix_models/mix_models_with_data保存 encoder 模型时会自动附加cls池化与归一化层,合并产物可直接用于SentenceTransformer检索流程。
总结
LM-Cocktail 以极低的工程成本,为"微调后通用性退化""新任务无训练可用模型""单任务模型无法多任务复用"三个常见痛点提供了统一的模型合并解法。配合 FlagEmbedding 生态中 bge 系列 Embedding 模型与 reranker,它既可以服务于检索场景(Embedding / Reranker 融合),也能覆盖生成场景(decoder LLM 融合)。三个公开函数分工明确:mix_models面向已知权重的手动合并、mix_models_with_data面向免训练的数据驱动自动加权、mix_models_by_layers面向大模型的内存友好逐层合并,开发者可以根据任务类型与资源约束灵活选用。
【免费下载链接】FlagEmbeddingRetrieval and Retrieval-augmented LLMs项目地址: https://gitcode.com/GitHub_Trending/fl/FlagEmbedding
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考