Autokeras StructuredDataClassifier选adam_weight_decay优化器时报错
解决Autokeras StructuredDataClassifier使用adam_weight_decay优化器时的类型转换错误
问题概述
使用Autokeras的StructuredDataClassifier做留一分组交叉验证训练,前4次试验均正常完成,但第5次试验自动选择adam_weight_decay作为优化器时,触发TensorFlow Graph执行错误,提示UNIMPLEMENTED: Cast string to float is not supported。
报错日志
Trial 5 Complete前日志略... 2022-12-11 16:22:23.607384: W tensorflow/core/framework/op_kernel.cc:1807] OP_REQUIRES failed at cast_op.cc:121 : UNIMPLEMENTED: Cast string to float is not supported tensorflow.python.framework.errors_impl.UnimplementedError: Graph execution error: ... Node: 'Cast_1' 2 root error(s) found. (0) UNIMPLEMENTED: Cast string to float is not supported [[{{node Cast_1}}]] (1) CANCELLED: Function was cancelled before it was started
复现代码
import tensorflow as tf import pandas as pd import numpy as np import autokeras as ak from sklearn.model_selection import LeaveOneGroupOut from sklearn.metrics import classification_report, confusion_matrix from sklearn.preprocessing import LabelEncoder data = pd.read_csv("p_feature_df.csv") y = data.pop('is_p') y = y.astype(np.int32) data.pop('idx') groups = data.pop('owner') data = data.astype(np.float32) X = data.to_numpy() lb = LabelEncoder() y = lb.fit_transform(y) logo = LeaveOneGroupOut() logo.get_n_splits(X,y,groups) results = [] models = [] for train_index, test_index in logo.split(X,y,groups): X_train, X_test = X[train_index], X[test_index] y_train, y_test = y[train_index], y[test_index] clf = ak.StructuredDataClassifier(overwrite=True) clf.fit(x=X_train, y=y_train, use_multiprocessing=True, workers=8, verbose=True) loss, acc = clf.evaluate(x=X_test, y=y_test, verbose=True) results.append( (loss, acc)) models.append(clf) print( (loss, acc) )
原因分析
尽管代码中已将特征数据转为float32,但adam_weight_decay优化器在Autokeras内部的处理逻辑存在特殊路径:
- 优化器初始化时,某些内部配置参数被误识别为字符串类型,触发了无效的类型转换
- 多进程训练模式下,数据传递过程中出现类型异常,恰好该优化器的试验逻辑未处理这种情况
解决方案
- 强制校验数据类型:转换为numpy数组后添加断言,确保无字符串残留:
X = data.to_numpy() assert np.issubdtype(X.dtype, np.floating), "训练数据存在非浮点类型特征" - 限制优化器可选范围:初始化
StructuredDataClassifier时指定允许的优化器,排除adam_weight_decay:clf = ak.StructuredDataClassifier( overwrite=True, optimizer=['adam', 'sgd', 'rmsprop'] # 仅保留稳定的优化器 ) - 禁用多进程训练:暂时关闭多进程,排查是否为进程间数据传递导致的问题:
clf.fit(x=X_train, y=y_train, verbose=True) # 移除use_multiprocessing和workers参数 - 升级依赖版本:更新Autokeras和TensorFlow到最新稳定版,修复已知的类型处理bug:
pip install --upgrade autokeras tensorflow
内容的提问来源于stack exchange,提问作者Shakuni
相关产品推荐
相关产品推荐

