微调Llama2模型:合并保存与文本嵌入生成的技术咨询
问题解答
1. 保存微调后的完整模型
你可以通过PEFT提供的权重合并方法,将基础模型与适配器权重合并后保存:
- 先将模型切换到评估模式,避免训练相关参数干扰:
model.eval() - 调用
merge_and_unload()方法合并权重,得到完整的模型实例:merged_model = model.merge_and_unload() - 最后用
save_pretrained()将合并后的模型保存到本地路径,同时可保存分词器方便后续使用:merged_model.save_pretrained("./full_finetuned_llama2") tokenizer.save_pretrained("./full_finetuned_llama2")
保存完成后,后续直接通过AutoModel.from_pretrained("./full_finetuned_llama2")即可加载完整模型,无需再加载PEFT适配器。
2. 基于已加载的模型生成文本嵌入
有两种方案适配LangChain的使用场景:
方案一:直接用加载的PeftModel自定义Embeddings类
由于LangChain的LlamaCppEmbeddings不支持传入模型实例,你可以继承Embeddings基类实现自定义嵌入类:
from langchain.embeddings.base import Embeddings from typing import List import torch class PeftLlamaEmbeddings(Embeddings): def __init__(self, model, tokenizer): self.model = model self.tokenizer = tokenizer self.device = model.device def embed_documents(self, texts: List[str]) -> List[List[float]]: embeddings = [] for text in texts: inputs = self.tokenizer( text, return_tensors="pt", truncation=True, max_length=512, padding=True ).to(self.device) with torch.no_grad(): outputs = self.model(**inputs) # 两种嵌入提取方式二选一:bos token输出 或 最后一层隐藏状态均值 embedding = outputs.last_hidden_state[:, 0, :].squeeze().cpu().numpy().tolist() # embedding = outputs.last_hidden_state.mean(dim=1).squeeze().cpu().numpy().tolist() embeddings.append(embedding) return embeddings def embed_query(self, text: str) -> List[float]: return self.embed_documents([text])[0]
实例化后即可在LangChain中使用:
# 假设你已加载好model(PeftModel实例)和tokenizer custom_embeddings = PeftLlamaEmbeddings(model, tokenizer) # 示例:用于构建向量数据库 from langchain.vectorstores import FAISS db = FAISS.from_texts(your_text_list, custom_embeddings)
方案二:合并模型后用HuggingFaceEmbeddings
如果已经保存了完整模型,直接使用LangChain内置的HuggingFaceEmbeddings加载即可:
from langchain.embeddings import HuggingFaceEmbeddings import torch embeddings = HuggingFaceEmbeddings( model_name="./full_finetuned_llama2", model_kwargs={"device": "cuda" if torch.cuda.is_available() else "cpu"}, encode_kwargs={"normalize_embeddings": True} # 可选,归一化可提升检索效果 )
该embeddings实例可直接用于LangChain的各类嵌入需求。
内容的提问来源于stack exchange,提问作者Hamid K
相关产品推荐
相关产品推荐

