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

如何在GPU上通过CUDA运行训练?解决Numba与Keras模型兼容报错

问题分析与解决方案

错误原因

你遇到的Cannot determine Numba type of <class 'keras.engine.training.Model'>错误,核心原因是Numba无法识别Keras模型对象的类型。@cuda.jit(nopython=True)要求函数内所有变量都能被Numba编译成本地机器码,但Keras的Model是高层框架对象,不属于Numba支持的原生类型(如numpy数组、基本数据类型)。

更关键的是:你完全不需要用Numba的CUDA装饰器来包裹整个训练流程——TensorFlow 1.x + Keras本身已经会自动检测并利用GPU进行训练,只要你的环境配置正确(CUDA、cuDNN版本匹配)。

修复步骤

1. 移除Numba的CUDA装饰器

直接删掉@cuda.jit(nopython=True)这一行,保留原本的训练函数逻辑即可。

2. 验证GPU是否被TensorFlow识别

在代码开头添加以下验证代码,确认TensorFlow能检测到GPU:

from tensorflow.python.client import device_lib
print(device_lib.list_local_devices())

如果输出中包含GPU设备信息,说明环境配置正确,Keras训练会自动跑在GPU上。

3. 修改后的训练函数代码

def training(nb_epoch,model,data_gen_train,data_gen_test,params,tr_loss,val_loss,sed_loss,doa_loss,sed_gt,doa_gt,epoch_metric_loss,unique_name,patience_cnt,conf_mat,best_conf_mat,best_metric):
    best_epoch = 0  # 补充初始化,避免未定义报错
    for epoch_cnt in range(nb_epoch):
        start = time.time()
        hist = model.fit_generator(
            generator=data_gen_train.generate(),
            steps_per_epoch=2 if params['quick_test'] else data_gen_train.get_total_batches_in_data(),
            validation_data=data_gen_test.generate(),
            validation_steps=2 if params['quick_test'] else data_gen_test.get_total_batches_in_data(),
            epochs=1,
            verbose=0
        )
        tr_loss[epoch_cnt] = hist.history.get('loss')[-1]
        val_loss[epoch_cnt] = hist.history.get('val_loss')[-1]

        pred = model.predict_generator(
            generator=data_gen_test.generate(),
            steps=2 if params['quick_test'] else data_gen_test.get_total_batches_in_data(),
            verbose=2
        )
        if params['mode'] == 'regr':
            sed_pred = evaluation_metrics.reshape_3Dto2D(pred[0]) > 0.5
            doa_pred = evaluation_metrics.reshape_3Dto2D(pred[1])

            sed_loss[epoch_cnt, :] = evaluation_metrics.compute_sed_scores(sed_pred, sed_gt, data_gen_test.nb_frames_1s())
            if params['azi_only']:
                doa_loss[epoch_cnt, :], conf_mat = evaluation_metrics.compute_doa_scores_regr_xy(doa_pred, doa_gt,
                                                                                                 sed_pred, sed_gt)
            else:
                doa_loss[epoch_cnt, :], conf_mat = evaluation_metrics.compute_doa_scores_regr_xyz(doa_pred, doa_gt,
                                                                                                  sed_pred, sed_gt)

            epoch_metric_loss[epoch_cnt] = np.mean([
                sed_loss[epoch_cnt, 0],
                1-sed_loss[epoch_cnt, 1],
                2*np.arcsin(doa_loss[epoch_cnt, 1]/2.0)/np.pi,
                1 - (doa_loss[epoch_cnt, 5] / float(doa_gt.shape[0]))]
            )
        plot_functions(unique_name, tr_loss, val_loss, sed_loss, doa_loss, epoch_metric_loss)

        patience_cnt += 1
        if epoch_metric_loss[epoch_cnt] < best_metric:
            best_metric = epoch_metric_loss[epoch_cnt]
            best_conf_mat = conf_mat
            best_epoch = epoch_cnt
            model.save('{}_model.h5'.format(unique_name))
            patience_cnt = 0

        print(
            'epoch_cnt: %d, time: %.2fs, tr_loss: %.2f, val_loss: %.2f, '
            'F1_overall: %.2f, ER_overall: %.2f, '
            'doa_error_gt: %.2f, doa_error_pred: %.2f, good_pks_ratio:%.2f, '
            'error_metric: %.2f, best_error_metric: %.2f, best_epoch : %d' %
            (
                epoch_cnt, time.time() - start, tr_loss[epoch_cnt], val_loss[epoch_cnt],
                sed_loss[epoch_cnt, 1], sed_loss[epoch_cnt, 0],
                doa_loss[epoch_cnt, 1], doa_loss[epoch_cnt, 2], doa_loss[epoch_cnt, 5] / float(sed_gt.shape[0]),
                epoch_metric_loss[epoch_cnt], best_metric, best_epoch
            )
        )
        if patience_cnt > params['patience']:
            break

额外说明

Numba的CUDA装饰器适用于底层数值计算函数(比如自定义的numpy数组运算),而Keras/TensorFlow的训练逻辑已经封装了GPU加速,包括模型前向/反向传播、数据并行等操作,不需要手动用Numba介入高层训练流程。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 02:30:44