GPT4All的prompt函数无报错停止返回值,寻求跳过异常迭代方案
解决方案
1. 为prompt调用添加超时控制(解决无响应阻塞)
GPT4All的prompt函数无响应时不会抛出异常,用线程池封装调用并设置超时,超时则终止任务返回None,避免阻塞整个流程:
from concurrent.futures import ThreadPoolExecutor, TimeoutError import threading # 自定义超时时间,可根据实际调整 PROMPT_TIMEOUT = 30 def single_chatgpt_offline(gpt, question): def prompt_task(): # GPT4All底层可能非线程安全,加锁避免冲突 with threading.Lock(): return gpt.prompt(question) try: print(f'\tAsking {question}') with ThreadPoolExecutor(max_workers=1) as executor: future = executor.submit(prompt_task) resp = future.result(timeout=PROMPT_TIMEOUT) print(f'\tGot the response: {resp}') return resp except TimeoutError: print(f'\tPrompt超时,跳过问题:{question}') # 超时后重置GPT实例,避免后续调用持续无响应 gpt.close() gpt = GPT4All() gpt.open() except Exception as err: print(f'\t调用出错:{str(err)}') return None
2. 修复Broken Pipe错误
管道中断通常是文件流意外中断导致的,在写入CSV时捕获该错误,确保流程继续:
# 替换原写入行的代码 try: csv_writer.writerow(line) except IOError as e: if e.errno == 32: print(f'第{idx+1}行写入失败:管道中断,跳过该行') # 可选:记录失败行到日志,方便后续补处理 with open('data/failed_rows.log', 'a', encoding='utf-8') as f: f.write(f'索引{idx}:{question}\n') else: print(f'写入错误:{str(e)}')
3. 补充优化细节
- 定义缺失的
NUM_DECIMAL变量(比如NUM_DECIMAL = 2) - 添加进度打印,每处理10条数据输出一次进度
- 确保异常退出时GPT4All实例被正确关闭
修改后的完整脚本
from nomic.gpt4all import GPT4All import csv from time import time from concurrent.futures import ThreadPoolExecutor, TimeoutError import threading CMU_CSV_PATH = 'data/cmu_qa.csv' PROMPT_TIMEOUT = 30 NUM_DECIMAL = 2 def single_chatgpt_offline(gpt, question): def prompt_task(): with threading.Lock(): return gpt.prompt(question) try: print(f'\tAsking {question}') with ThreadPoolExecutor(max_workers=1) as executor: future = executor.submit(prompt_task) resp = future.result(timeout=PROMPT_TIMEOUT) print(f'\tGot the response: {resp}') return resp except TimeoutError: print(f'\tPrompt超时,跳过问题:{question}') gpt.close() gpt = GPT4All() gpt.open() except Exception as err: print(f'\t调用出错:{str(err)}') return None def find_answers_and_write_csv(): gpt = None try: gpt = GPT4All() gpt.open() with open(CMU_CSV_PATH, 'r', encoding='utf-8') as csv_src: csv_reader = csv.DictReader(csv_src) with open('data/cmu_qa_answers.csv', 'w', encoding='utf-8', newline='') as csv_target: fields = ['question', 'answer', 'title', 'bard', 'bard_time', 'gpt', 'gpt_time'] csv_writer = csv.DictWriter(csv_target, fieldnames=fields) csv_writer.writeheader() for idx, line in enumerate(csv_reader): if idx % 10 == 0: print(f'已处理{idx}条数据') question = line['question'] line['title'] = line['title'].replace('_', ' ') duration_gpt_start = time() resp_gpt = single_chatgpt_offline(gpt, question) if resp_gpt and '\n' in resp_gpt: resp_gpt = resp_gpt.replace('\n', ' ') duration_gpt_end = time() line['gpt'] = resp_gpt line['gpt_time'] = round((duration_gpt_end - duration_gpt_start), NUM_DECIMAL) line["bard"] = None line["bard_time"] = None try: csv_writer.writerow(line) except IOError as e: if e.errno == 32: print(f'第{idx+1}行写入失败:管道中断,跳过该行') with open('data/failed_rows.log', 'a', encoding='utf-8') as f: f.write(f'索引{idx}:{question}\n') else: print(f'写入错误:{str(e)}') finally: if gpt is not None: gpt.close() if __name__ == '__main__': find_answers_and_write_csv()
内容的提问来源于stack exchange,提问作者talha06
相关产品推荐
相关产品推荐

