You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Python并行处理中JSON存储效率优化及代码改进问询

统一存储共享JSON的最优方案

最可靠且高效的实现方式是利用multiprocessing.Manager创建进程安全的共享字典,让所有子进程将结果存入该字典,最后由主进程统一写入单个JSON文件。这种方式避免了多进程直接写文件的冲突问题,同时减少了频繁文件IO的开销。

修改步骤及代码示例

  1. 导入依赖模块

    import functools
    import multiprocessing
    
  2. 修改call_api函数
    移除单个文件保存逻辑,将结果存入共享字典,同时增加异常处理避免进程崩溃:

    def call_api(example, shared_results):
        try:
            # 发送API请求
            response = requests.post(MIXSTRAL_URL, json=example['input'])
            # 检查响应状态
            if check_status(response.status_code):
                query_pred = extraction_pipe(response.text)
                # 构造结果条目
                result_entry = {
                    'question': example['question'],
                    'query_gold': example['query_gold'],
                    'query_pred': {str(i): query for i, query in enumerate(query_pred)}
                }
                # 存入共享字典
                shared_results[example['id']] = result_entry
        except Exception as e:
            # 记录错误信息,可根据需求写入日志文件
            print(f"处理ID {example['id']} 失败: {str(e)}")
            shared_results[example['id']] = {'error': str(e)}
    
  3. 修改主函数mixtral_prompting_similarity_parallel
    创建共享字典,用functools.partial包装call_api传入共享字典,最后统一写入JSON:

    def mixtral_prompting_similarity_parallel(test_ds, train_ds, save_path, prompt, similarities, train_data_rel=None, test_data_rel=None):
        type_example = 'ent_rel_unique' if train_data_rel and test_data_rel else ''
        all_input = {"id": [], "input": [], "question": [], "query_gold": []}
    
        # 构建输入数据(移除原save_path字段,不再需要)
        for id, question_dict in test_ds.items():
            input_config = config.copy()
            prompt_examples = ""
            for sim_id in similarities[id]:
                q_dict = train_data_rel[sim_id] if type_example else train_ds[sim_id]
                prompt_examples += create_example(q_dict, type_example)
    
            type_val = type_example if type_example else ''
            q_dict = test_data_rel[id] if type_example else question_dict
    
            input = create_mixstral_input(q_dict, type_val)
            input_config['text'] = prompt + prompt_examples + input
    
            all_input["id"].append(id)
            all_input["input"].append(input_config)
            all_input["question"].append(question_dict['question'])
            all_input["query_gold"].append(question_dict['query'])
    
        df = Dataset.from_dict(all_input)
        # 使用Manager创建共享字典
        with multiprocessing.Manager() as manager:
            shared_results = manager.dict()
            # 包装call_api,传入共享字典参数
            wrapped_call_api = functools.partial(call_api, shared_results=shared_results)
            # 动态设置进程数(推荐用CPU核心数或API并发限制)
            with multiprocessing.Pool(processes=multiprocessing.cpu_count()) as pool:
                pool.map(wrapped_call_api, df)
            # 主进程统一写入JSON文件
            final_results = dict(shared_results)
            save_json(os.path.join(save_path, 'all_predictions.json'), final_results)
    

其他效率优化建议

1. 复用HTTP连接减少开销

每次调用requests.post都会新建TCP连接,可在每个子进程中初始化requests.Session复用连接:

def init_worker():
    global session
    session = requests.Session()

# 在创建Pool时指定初始化函数
with multiprocessing.Pool(processes=..., initializer=init_worker) as pool:
    # ... 执行任务 ...

# 修改call_api中的请求方式
response = session.post(MIXSTRAL_URL, json=example['input'])

2. 合理设置进程数

避免硬编码进程数(如原代码中的16),建议根据CPU核心数或API的QPS限制动态调整:

# 使用CPU核心数
processes = multiprocessing.cpu_count()
# 若API有并发限制,取较小值
processes = min(processes, API_MAX_CONCURRENCY)

3. 替换多进程为异步IO(IO密集型场景更高效)

调用API属于IO密集型任务,用aiohttp+asyncio替代多进程,可减少进程切换开销:

import asyncio
import aiohttp

async def async_call_api(session, example):
    try:
        async with session.post(MIXSTRAL_URL, json=example['input']) as response:
            if check_status(response.status):
                text = await response.text()
                query_pred = extraction_pipe(text)
                return (example['id'], {
                    'question': example['question'],
                    'query_gold': example['query_gold'],
                    'query_pred': {str(i): query for i, query in enumerate(query_pred)}
                })
    except Exception as e:
        return (example['id'], {'error': str(e)})

async def main():
    async with aiohttp.ClientSession() as session:
        tasks = [async_call_api(session, example) for example in df]
        results = await asyncio.gather(*tasks)
        final_results = dict(results)
        save_json(os.path.join(save_path, 'all_predictions.json'), final_results)

# 运行异步任务
asyncio.run(main())

4. 增强异常与错误处理

添加详细的异常捕获(如连接超时、JSON解析失败等),并将错误信息存入结果,方便后续排查问题。

5. 优化数据预处理

如果输入数据构建过程耗时,可考虑并行预处理(如用multiprocessing.Pool并行生成input_config),减少主线程的等待时间。


内容的提问来源于stack exchange,提问作者jacopodabramo

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 13:28:25