TensorFlow 1.4.0中DNNLinearCombinedClassifier报错:标签ID需≥0
解决TensorFlow 1.4.0中DNNLinearCombinedClassifier的Label IDs错误
这个InvalidArgumentError的核心原因很明确:你的字符串标签没有被正确转换为非负整数ID,当TensorFlow遇到不在你定义的LABELS列表中的标签值时,会默认返回-1,这就触发了"Label IDs must >= 0"的断言检查。
下面是具体的修复步骤和代码调整:
1. 确保标签与词汇表完全匹配
首先检查你的数据集里的section_code字段值,是否和LABELS列表中的每个条目完全一致(包括开头的空格,比如' TS4CFS6')。如果数据里的标签少了空格或者有拼写错误,就会被识别为未知标签,生成无效的-1 ID。
2. 在输入函数中添加标签转换逻辑
你当前的代码缺少了关键的标签预处理步骤——需要把字符串标签映射为从0开始的整数ID,这是DNNLinearCombinedClassifier要求的标签格式。
添加完整的输入函数,并在其中处理标签转换:
def input_fn(data_file, num_epochs=1, shuffle=True, batch_size=64): def parse_csv(line): # 解析CSV行到特征字典 columns = tf.decode_csv(line, record_defaults=CSV_COLUMN_DEFAULTS) features = dict(zip(CSV_COLUMNS, columns)) # 移除无用列 for col in UNUSED_COLUMNS: features.pop(col) # 提取标签并转换为整数ID label_str = features.pop(LABEL_COLUMN) # 使用词汇表将字符串标签映射为整数,未知标签会返回-1(这里要确保数据中没有未知标签) label_id = tf.contrib.lookup.string_to_index( label_str, vocabulary_list=LABELS, default_value=-1 ) # 添加断言,提前捕获无效标签 label_id = tf.assert_non_negative( label_id, message="发现无效标签:不在LABELS列表中" ) return features, label_id # 构建数据集 dataset = tf.data.TextLineDataset(data_file) if shuffle: dataset = dataset.shuffle(buffer_size=10000) # 并行解析CSV dataset = dataset.map(parse_csv, num_parallel_calls=multiprocessing.cpu_count()) dataset = dataset.repeat(num_epochs) dataset = dataset.batch(batch_size) # 获取特征和标签 iterator = dataset.make_one_shot_iterator() features, labels = iterator.get_next() return features, labels
3. 初始化字符串查找表(TensorFlow 1.x必需)
在TensorFlow 1.x中,tf.contrib.lookup的词汇表需要手动初始化。在启动训练前,添加以下代码:
# 创建会话并初始化查找表 sess = tf.Session() sess.run(tf.tables_initializer()) # 或者在使用Estimator时,通过hooks初始化 from tensorflow.python.training.session_run_hook import SessionRunHook class InitLookupTablesHook(SessionRunHook): def after_create_session(self, session, coord): session.run(tf.tables_initializer()) # 训练时添加hook estimator.train( input_fn=lambda: input_fn("train_data.csv"), hooks=[InitLookupTablesHook()] )
4. 额外检查:确认LABELS覆盖所有标签值
如果你的数据中存在LABELS列表之外的标签,即使添加了转换逻辑,还是会生成-1 ID导致报错。建议先统计数据中的所有section_code值,确保它们都在LABELS列表里。可以用Python的pandas提前检查:
import pandas as pd df = pd.read_csv("your_data.csv") unique_labels = df['section_code'].unique() # 检查是否有不在LABELS中的标签 missing_labels = [label for label in unique_labels if label not in LABELS] if missing_labels: print("发现未在LABELS中定义的标签:", missing_labels)
这样调整后,你的标签会被正确转换为非负整数ID,就能解决这个断言错误了。
内容的提问来源于stack exchange,提问作者Binh Pham
相关产品推荐
相关产品推荐

