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

Python中DataFrame真值歧义报错:ANN训练时的函数使用问题

ANN训练中argmin调用触发ValueError的解决方法

问题重现

训练多种人工神经网络(ANN)时,执行analyze_results函数中的代码行:

optimized_params = rainfall_data.loc[rainfall_data.RMSE.argmin]

触发ValueError,报错信息:

ValueError: The truth value of a DataFrame is ambiguous. Use a.empty, a.bool(), a.item(), a.any() or a.all().

原函数代码:

def analyze_results(rainfall_data, test_rainfall_data, name, flag=False):
    optimized_params = rainfall_data.loc[rainfall_data.RMSE.argmin]
    future_steps = optimized_params.future_steps
    forecast_values = optimized_params[-1*int(future_steps):]
    y_true = test_rainfall_data.iloc[:int(future_steps)]
    forecast_values.index = y_true.index
    
    print('=== Best parameters of ' + name + ' ===\n')
    if (name == 'FNN' or name == 'LSTM'):
        model = create_NN(optimized_params.look_back, 
                          optimized_params.hidden_nodes, 
                          optimized_params.output_nodes)
        print('Input nodes(p): ' + str(optimized_params.look_back))
        print('Hidden nodes: ' + str(optimized_params.hidden_nodes))
        print('Output nodes: ' + str(optimized_params.output_nodes))
    elif (name == 'TLNN'):
        model = create_NN(len(optimized_params.look_back_lags), 
                          optimized_params.hidden_nodes, 
                          optimized_params.output_nodes)
        s = ''
        for i in optimized_params.look_back_lags:
            s = s+' '+str(i)
        print('Look back lags: ' + s)
        print('Hidden nodes: ' + str(optimized_params.hidden_nodes))
        print('Output nodes: ' + str(optimized_params.output_nodes))
    elif (name == 'SANN'):
        model = create_NN(optimized_params.seasonal_period, 
                          optimized_params.hidden_nodes, 
                          optimized_params.seasonal_period)
        print('Input nodes(s): ' + str(optimized_params.seasonal_period))
        print('Hidden nodes: ' + str(optimized_params.hidden_nodes))
        print('Output nodes: ' + str(optimized_params.seasonal_period))
        
    print('Number of epochs: ' + str(optimized_params.epochs))
    print('Batch size: ' + str(optimized_params.batch_size))
    print('Number of future steps forecasted: ' + str(optimized_params.future_steps))
    print('Mean Squared Error(MSE): ' + str(optimized_params.MSE))
    print('Mean Absolute Error(MAE): ' + str(optimized_params.MAE))
    print('Root Mean Squared Error(RMSE): ' + str(optimized_params.RMSE))
    print('\n\n')

完整报错栈:

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
~\AppData\Local\Temp\ipykernel_12964\1785142711.py in ?()
      9 
     10 # look_back, hidden_nodes, output_nodes, epochs, batch_size, future_steps
     11 parameters_LSTM = [[1,2,3,4,5,6,7,8,9,10,11,12,13], [3,4,5,6], [1], [300], [20], [future_steps]]
     12 
---> 13 RMSE_info = compare_ANN_methods(rainfall_data, test_rainfall_data, scaler, parameters_FNN, parameters_TLNN, parameters_SANN, parameters_LSTM, future_steps)

~\AppData\Local\Temp\ipykernel_12964\2478982653.py in ?(rainfall_data, test_rainfall_data, scaler, parameters_FNN, parameters_TLNN, parameters_SANN, parameters_LSTM, future_steps)
      1 def compare_ANN_methods(rainfall_data, test_rainfall_data, scaler, parameters_FNN, parameters_TLNN, parameters_SANN, parameters_LSTM, future_steps):
      2 
      3     information_FNN_df = get_accuracies_FNN(rainfall_data, test_rainfall_data, parameters_FNN, scaler)
----> 4     optimized_params_FNN = analyze_results(information_FNN_df, test_rainfall_data, 'FNN')
      5 
      6     information_TLNN_df = get_accuracies_TLNN(rainfall_data, test_rainfall_data, parameters_TLNN, scaler)
      7     optimized_params_TLNN = analyze_results(information_TLNN_df, test_rainfall_data, 'TLNN')

~\AppData\Local\Temp\ipykernel_12964\4019196368.py in ?(rainfall_data, test_rainfall_data, name, flag)
      1 def analyze_results(rainfall_data, test_rainfall_data, name, flag=False):
----> 2     optimized_params = rainfall_data.loc[(rainfall_data.RMSE.argmin)]
      3     future_steps = optimized_params.future_steps
      4     forecast_values = optimized_params[-1*int(future_steps):]
      5     y_true = test_rainfall_data.iloc[:int(future_steps)]

~\anaconda3\Lib\site-packages\pandas\core\indexing.py in ?(self, key)
   1185         else:
   1186             # we by definition only have the 0th axis
   1187             axis = self.axis or 0
   1188 
-> 1189             maybe_callable = com.apply_if_callable(key, self.obj)
   1190             maybe_callable = self._check_deprecated_callable_usage(key, maybe_callable)
   1191             return self._getitem_axis(maybe_callable, axis=axis)

~\anaconda3\Lib\site-packages\pandas\core\common.py in ?(maybe_callable, obj, **kwargs)
    380     obj : NDFrame
    381     **kwargs
    382     """
    383     if callable(maybe_callable):
--> 384         return maybe_callable(obj, **kwargs)
    385 
    386     return maybe_callable

~\anaconda3\Lib\site-packages\pandas\core\base.py in ?(self, axis, skipna, *args, **kwargs)
    765     def argmin(
    766         self, axis: AxisInt | None = None, skipna: bool = True, *args, **kwargs
    767     ) -> int:
    768         delegate = self._values
--> 769         nv.validate_minmax_axis(axis)
    770         skipna = nv.validate_argmin_with_skipna(skipna, args, kwargs)
    771 
    772         if isinstance(delegate, ExtensionArray):

~\anaconda3\Lib\site-packages\pandas\compat\numpy\function.py in ?(axis, ndim)
    395     ValueError
    396     """
    397     if axis is None:
    398         return
--> 399     if axis >= ndim or (axis < 0 and ndim + axis < 0):
    400         raise ValueError(f"`axis` must be fewer than the number of dimensions ({ndim})")

~\anaconda3\Lib\site-packages\pandas\core\generic.py in ?(self)
   1575     @final
   1576     def __nonzero__(self) -> NoReturn:
-> 1577         raise ValueError(
   1578             f"The truth value of a {type(self).__name__} is ambiguous. "
   1579             "Use a.empty, a.bool(), a.item(), a.any() or a.all()."
   1580         )

ValueError: The truth value of a DataFrame is ambiguous. Use a.empty, a.bool(), a.item(), a.any() or a.all().

报错原因

核心问题是代码中调用argmin方法时遗漏了括号:

  • rainfall_data.RMSE.argmin是方法对象本身,而非方法调用后的结果
  • pandas的loc索引器会将可调用对象传入当前DataFrame作为参数执行,导致argmin方法被调用时,axis参数被错误地传入了整个rainfall_data DataFrame
  • 在后续的轴验证逻辑中,尝试对DataFrame进行布尔比较(axis >= ndim),触发了DataFrame布尔值判断的歧义错误

解决方案

将代码中的rainfall_data.RMSE.argmin修改为rainfall_data.RMSE.argmin(),确保调用方法并返回RMSE列最小值对应的索引整数:

# 修改前
optimized_params = rainfall_data.loc[rainfall_data.RMSE.argmin]

# 修改后
optimized_params = rainfall_data.loc[rainfall_data.RMSE.argmin()]

额外验证

如果rainfall_data.RMSE是多列DataFrame(而非一维Series),需先确认目标RMSE列的名称,再指定单列调用argmin(),例如:

# 明确指定目标RMSE列
optimized_params = rainfall_data.loc[rainfall_data['RMSE'].argmin()]

内容的提问来源于stack exchange,提问作者Vivek sreekumar menon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 06:14:50