llama_index 0.7.1导入CustomLLM失败,求自定义LLM适配方案
解决llama-index 0.7.1无法导入CustomLLM的问题
问题原因
你遇到的导入错误是因为llama-index的模块结构在0.8.x版本后才做了调整,CustomLLM、CompletionResponse、LLMMetadata这些类是0.8.x及以上版本才移到llama_index.llms模块下的,而你使用的0.7.1版本还沿用旧的架构,所以官方新版本的示例代码无法直接在0.7.1中运行。
可行解决方案
方案一:升级llama-index到0.8.x及以上版本
如果没有版本依赖限制,直接升级到最新版本即可兼容官方示例代码:
pip install --upgrade llama-index
升级后原示例代码的导入和逻辑都能正常运行。
方案二:适配0.7.1版本的自定义LLM写法
如果必须保留0.7.1版本,需要修改代码以适配旧版的LLMPredictor架构,具体修改如下:
替换原代码中的导入和自定义LLM类部分:
import torch from transformers import pipeline from typing import Optional, List, Mapping, Any from llama_index import ( ServiceContext, SimpleDirectoryReader, LangchainEmbedding, ListIndex ) # 替换原llms模块的导入,改用0.7.1的LLMPredictor基类 from llama_index.llm_predictor.base import BaseLLMPredictor from llama_index.prompts.base import Prompt # set context window size context_window = 2048 # set number of output tokens num_output = 256 # store the pipeline/model outisde of the LLM class to avoid memory issues model_name = "facebook/opt-iml-max-30b" pipeline = pipeline("text-generation", model=model_name, device="cuda:0", model_kwargs={"torch_dtype":torch.bfloat16}) # 自定义LLMPredictor替代原CustomLLM class CustomLLMPredictor(BaseLLMPredictor): def __init__(self, pipeline, context_window: int = 2048, num_output: int = 256): self.pipeline = pipeline self.context_window = context_window self.num_output = num_output def predict(self, prompt: Prompt, **kwargs: Any) -> str: prompt_str = prompt.get_text() response = self.pipeline(prompt_str, max_new_tokens=self.num_output)[0]["generated_text"] # 返回新增生成的内容 return response[len(prompt_str):] # 批量预测方法可选实现,这里做简单处理 def predict_batch(self, prompts: List[Prompt], **kwargs: Any) -> List[str]: return [self.predict(prompt) for prompt in prompts] # 初始化自定义predictor llm_predictor = CustomLLMPredictor(pipeline) service_context = ServiceContext.from_defaults( llm_predictor=llm_predictor, # 替换原llm参数为llm_predictor context_window=context_window, num_output=num_output ) # 后续加载数据、构建索引、查询的代码保持不变 documents = SimpleDirectoryReader('./data').load_data() index = ListIndex.from_documents(documents, service_context=service_context) query_engine = index.as_query_engine() response = query_engine.query("<query_text>") print(response)
修改说明:
- 用
BaseLLMPredictor替代CustomLLM作为基类,这是0.7.x版本自定义LLM的标准方式 - 重写
predict方法实现文本生成逻辑,与原示例的complete方法逻辑一致 - 创建
ServiceContext时传入llm_predictor参数,而非新版的llm参数
内容的提问来源于stack exchange,提问作者Peyman
相关产品推荐
相关产品推荐

