Python如何向函数传入列表全部元素 批量抽取文本关系三元组
问题根因
原代码仅处理首条文本的核心原因有两个:
- 传入pipeline时主动取了返回结果的第0个元素
[0]["generated_token_ids"],直接丢弃了其余文本的生成结果 - 没有配置pipeline的批量推理逻辑,传入多文本集合时没有按规则处理所有生成的token序列
for循环完全可以实现全量文本遍历处理,同时transformers的pipeline本身支持原生批量推理,处理效率比逐行for循环更高
可直接运行的完整实现
注意你当前定义的rez_dictionary是set集合类型,遍历顺序随机、会自动去重,建议先转为有序列表存储待处理文本,以下代码完全适配你已有的三元组解析函数,可直接输出结构化结果:
from transformers import pipeline # 1. 整理待处理文本为有序列表 text_list = [ 'Decent Little Reader, Poor Tablet', 'Ok For What It Is', 'Too Heavy and Poor weld quality,', 'difficult mount', 'just got it installed' ] # 2. 加载三元组抽取pipeline triplet_extractor = pipeline( 'text2text-generation', model='Babelscape/rebel-large', tokenizer='Babelscape/rebel-large', device_map="auto" # 有GPU时会自动调用加速推理 ) # 3. 原有三元组解析函数(无需修改) def extract_triplets(text): triplets = [] relation, subject, relation, object_ = '', '', '', '' text = text.strip() current = 'x' for token in text.replace("<s>", "").replace("<pad>", "").replace("</s>", "").split(): if token == "<triplet>": current = 't' if relation != '': triplets.append({'head': subject.strip(), 'type': relation.strip(),'tail': object_.strip()}) relation = '' subject = '' elif token == "<subj>": current = 's' if relation != '': triplets.append({'head': subject.strip(), 'type': relation.strip(),'tail': object_.strip()}) object_ = '' elif token == "<obj>": current = 'o' relation = '' else: if current == 't': subject += ' ' + token elif current == 's': object_ += ' ' + token elif current == 'o': relation += ' ' + token if subject != '' and relation != '' and object_ != '': triplets.append({'head': subject.strip(), 'type': relation.strip(),'tail': object_.strip()}) return triplets # -------------------------- # 方式1:原生批量推理(推荐,处理效率更高) # -------------------------- # 批量生成所有文本的token序列,不再仅取第0位结果 generated_outputs = triplet_extractor( text_list, return_tensors=True, return_text=False, batch_size=4 # 可根据显存/内存大小调整批次大小 ) # 批量解码所有生成结果 generated_texts = triplet_extractor.tokenizer.batch_decode( [item["generated_token_ids"] for item in generated_outputs] ) # 遍历解码文本,解析得到全量三元组 all_triplets = [] for text in generated_texts: all_triplets.extend(extract_triplets(text)) print(all_triplets) # -------------------------- # 方式2:for循环逐行处理(逻辑直观,适合小数据量场景) # -------------------------- all_triplets_loop = [] for single_text in text_list: # 单条文本传入生成结果 single_output = triplet_extractor( single_text, return_tensors=True, return_text=False )[0]["generated_token_ids"] # 解码单条结果 single_decoded_text = triplet_extractor.tokenizer.decode(single_output) # 解析三元组加入结果集 all_triplets_loop.extend(extract_triplets(single_decoded_text)) print(all_triplets_loop)
关键注意点
- 不要用无键值对的
{}集合存储待处理文本,集合自带去重、无序特性,可能导致文本丢失、顺序错乱,统一用列表存储即可 - 批量推理时可根据硬件配置调整
batch_size参数,显存/内存越大,可设置的数值越高,处理速度越快 - 两种处理方式输出的三元组结构完全一致,都能直接适配原有解析函数的输入要求,不需要修改解析逻辑
内容的提问来源于stack exchange,提问作者Timskouten
相关产品推荐
相关产品推荐

