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

Python中使用N-Gram语言模型时map函数的多进程错误处理

问题描述

我希望通过N-Gram模型提升Speech2Text模型的准确率,因此使用以下代码对整个数据集应用处理函数:

result = dataset.map(predict, batch_size=5, num_proc=int(os.environ.get('cpu_core')))

其中cpu_core设置为8。predict函数代码如下:

def predict(batch):
     batch["predicted"] = processor.batch_decode(np.array(batch["logits"])).text[0]
     print(batch["predicted"])
     return batch

该代码位于while True循环的try块中,当程序遇到多进程错误时会陷入循环无法退出,完整代码如下:

while True:

    try:
        dataset = dataset.map(speech_file_to_array_fn)

        # 启用N-Gram模型时的处理
        if os.environ.get('active_ngram') == '1':
            dataset = dataset.map(predict_model)

        print("\nN-Gram started\n")
        result = dataset.map(predict, batch_size=5, num_proc=int(os.environ.get('cpu_core'))) # 触发错误的代码行

    except KeyboardInterrupt:
        print('interrupted!')
        break  

    except:
        pass

运行时出现多进程错误,错误信息显示多个ForkPoolWorker进程触发KeyboardInterrupt,当前环境为Python 3.8.10和Ubuntu 20.04.4,请问如何处理该多进程错误?

解决方案

针对多进程触发KeyboardInterrupt导致循环无法退出的问题,可从以下几个方向修复:

1. 优化异常捕获逻辑,避免无限循环

当前except:会静默忽略所有异常,包括多进程错误,导致程序无限重试。需要明确捕获异常并添加退出控制:

import traceback

retry_count = 0
max_retries = 3

while True:
    try:
        # 原有数据集处理流程
        dataset = dataset.map(speech_file_to_array_fn)
        if os.environ.get('active_ngram') == '1':
            dataset = dataset.map(predict_model)
        print("\nN-Gram started\n")
        result = dataset.map(predict, batch_size=5, num_proc=int(os.environ.get('cpu_core')))
        
        # 处理完成后主动退出循环
        break

    except KeyboardInterrupt:
        print('interrupted!')
        break  
    except Exception as e:
        # 打印错误详情以便排查
        print(f"错误详情: {str(e)}")
        traceback.print_exc()
        
        # 限制重试次数,避免无限循环
        retry_count += 1
        if retry_count >= max_retries:
            print(f"已重试{max_retries}次,程序退出")
            break
        print(f"第{retry_count}次重试...")

2. 修复predict函数的潜在逻辑错误

processor.batch_decode的返回格式可能不符合预期,.text[0]的写法在多进程环境下容易触发索引错误。先确认解码逻辑正确性:

def predict(batch):
    decoded_results = processor.batch_decode(np.array(batch["logits"]))
    # 根据实际返回格式调整,若返回列表则直接赋值,否则取对应属性
    batch["predicted"] = decoded_results if isinstance(decoded_results, list) else decoded_results.text
    print(batch["predicted"])
    return batch

3. 调整多进程参数,避免资源过载

num_proc=8可能导致CPU资源耗尽,引发进程间异常。可降低进程数或根据系统核心数动态调整:

import multiprocessing

# 取环境变量设置值和实际CPU核心数的较小值
num_proc = min(int(os.environ.get('cpu_core', 4)), multiprocessing.cpu_count())
result = dataset.map(predict, batch_size=5, num_proc=num_proc)

4. 先以单进程模式排查基础错误

先将num_proc设为1,确认predict函数本身无逻辑错误后,再逐步启用多进程:

result = dataset.map(predict, batch_size=5, num_proc=1)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 08:00:29