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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 16:25:26