使用Promptify做多标签文本分类遇Pipeline.fit()参数错误求助
问题原因与解决方案
错误根源
你调用Pipeline.fit()时出现got multiple values for argument 'text_input'的原因是:
- 创建
Pipeline实例时已经传入了绑定好模板的prompter对象(模板为multilabel_classification.jinja) - 调用
fit()时又将模板文件名作为第一个位置参数传入,而Pipeline.fit()的第一个位置参数默认对应text_input,这就导致你既通过位置参数传递了一个值给text_input,又通过关键字参数text_input=sent传递了另一个值,触发参数重复绑定错误。
修正后的代码
model = OpenAI(api_key) prompter = Prompter('multilabel_classification.jinja') pipe = Pipeline(prompter, model) classes = ['Medicine','Oncology','Metastasis','Breast cancer','Lung cancer','Cerebrospinal fluid','Tumor microenvironment','Single-cell RNA sequencing','Idiopathic intracranial hypertension'] sent = "The patient is a 93-year-old female with a medical history of chronic right hip pain, osteoporosis, hypertension, depression, and chronic atrial fibrillation admitted for evaluation and management of severe nausea and vomiting and urinary tract infection" # 移除多余的模板文件名参数,仅传递必要的关键字参数 result = pipe.fit( text_input=sent, n_output_labels=len(classes), domain='Clinical', labels=classes ) print(eval(result['text']))
补充说明
如果后续需要动态切换分类模板,不要直接在fit()中传递模板文件名,而是重新创建对应的Prompter实例并初始化新的Pipeline,或者参考Promptify官方文档中关于动态模板切换的方法。
内容的提问来源于stack exchange,提问作者Jasym
相关产品推荐
相关产品推荐

