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

TensorFlow model.fit报错:Invalid dtype: object 问题求助

胸部X光多标签分类模型训练时触发ValueError: Invalid dtype: object错误

在Kaggle上训练胸部X光多标签分类模型,使用ImageDataGenerator从DataFrame加载数据,自定义process_labels函数处理多标签编码,但调用model.fit时触发ValueError: Invalid dtype: object错误。尝试在数据导入、生成器创建、模型拟合前等多个阶段编码均未解决,以下是完整代码及报错堆栈:

完整代码

# Define transformations
train_datagen = ImageDataGenerator(
    rescale=1./255,  # Normalize pixel values
    rotation_range=20,  # Random rotations
    width_shift_range=0.2,  # Random horizontal shifts
    height_shift_range=0.2,  # Random vertical shifts
    shear_range=0.2,  # Random shearing transformations
    zoom_range=0.2,  # Random zoom
    horizontal_flip=True,  # Random horizontal flipping
    fill_mode='nearest'  # How to fill in new pixels after transformations
)

# Define datasets based on gender and view position
datasets = {
    'male_pa': df[(df['Patient Gender'] == 'M') & (df['View Position'] == 'PA')],
    'male_ap': df[(df['Patient Gender'] == 'M') & (df['View Position'] == 'AP')],
    'female_pa': df[(df['Patient Gender'] == 'F') & (df['View Position'] == 'PA')],
    'female_ap': df[(df['Patient Gender'] == 'F') & (df['View Position'] == 'AP')],
}

# Function to process labels
def process_labels(label):
    labels = label.split('|')
    encoded_labels = np.zeros(num_classes, dtype=int)  # Create a zero array
    for idx, disease in enumerate(all_diseases):
        if disease in labels:
            encoded_labels[idx] = 1
    return encoded_labels
# Get all unique diseases
all_diseases = set('|'.join(df['Finding Labels']).split('|'))
num_classes = len(all_diseases)  # Number of classes

# Create and process generators for each dataset
train_datasets = {}
validation_datasets = {}
for name, subset_df in datasets.items():


    # Determine all possible image folders based on the data
    image_folders = [os.path.join(data_dir, f"images_{i:03d}", "images")
                     for i in range(1, 13)]
    print(image_folders)

    def get_image_path(image_index):
        for folder in image_folders:
            full_path = os.path.join(folder, image_index)
            if os.path.exists(full_path):
                return full_path
        # Debugging print for missing images
        return None

    # Construct the 'image_path' column, searching for files across folders
    subset_df = subset_df.copy()  # Make explicit copies to ensure modification
    subset_df.loc[:, 'image_path'] = subset_df['Image Index'].apply(get_image_path)

    # Filter out missing images (if needed)
    subset_df = subset_df[subset_df['image_path'].notnull()]

    # Split data into training and test sets
    train_df, test_df = train_test_split(subset_df, test_size=0.2, random_state=42)

    # Training dataset generator
    train_dataset_generator = train_datagen.flow_from_dataframe(
        dataframe=train_df,
        directory=None,
        target_size=(224, 224),
        batch_size=32,
        class_mode='raw',
        x_col="image_path",
        y_col="Finding Labels",
        y_col_preprocessor=process_labels,
        shuffle=False,
        interpolation="nearest",
        validate_filenames=True,
        dtype='int32'  # <-- Set the data type for the labels
    )

    train_datasets[name] = train_dataset_generator
    print("hi")
    # Test dataset generator
    test_datagen = ImageDataGenerator(rescale=1./255)  # No augmentation for test
    test_dataset_generator = test_datagen.flow_from_dataframe(
        dataframe=test_df,
        directory=None,
        target_size=(224, 224),
        batch_size=32,
        class_mode='raw',
        x_col="image_path",
        y_col="Finding Labels",
        y_col_preprocessor=process_labels,
        shuffle=False,
        interpolation="nearest",
        validate_filenames=True,
        dtype='int32'  # <-- Set the data type for the labels
    )

validation_datasets[name] = test_dataset_generator

# Print the number of classes
print("Number of classes:", num_classes)

for dataset_name, train_dataset in train_datasets.items():

    print(f"Training model for dataset: {dataset_name}")

    # Load the ResNet50 model without the top classification layer
    base_model = ResNet50(weights='imagenet', include_top=False, input_shape=(224, 224, 3))

    # Freeze the base model's layers
    for layer in base_model.layers:
        layer.trainable = False

    # Add new classification layers on top of the base model
    x = base_model.output
    x = tf.keras.layers.GlobalAveragePooling2D()(x)
    x = tf.keras.layers.Dense(1024, activation='relu')(x)
    predictions = tf.keras.layers.Dense(num_classes, activation='sigmoid')(x)

    # Create the final model
    model = tf.keras.Model(inputs=base_model.input, outputs=predictions)

    # Compile the model
    model.compile(optimizer='adam',
                  loss='binary_crossentropy',
                  metrics=['accuracy'])

    # Train the model
    epochs = 10
    model.fit(train_dataset,
              epochs=epochs,
              verbose=1,
              validation_data=train_dataset,
              validation_freq=1)  # Validate after each epoch

    # Save the trained model
    model.save(f'chest_xray_model_{dataset_name}.h5')

报错堆栈

---------------------------------------------------------------------------
ValueError                                Traceback (most recent call last)
Cell In[91], line 28
     26 # Train the model
     27 epochs = 10
---> 28 model.fit(train_dataset,
     29           epochs=epochs,
     30           verbose=1,
     31           validation_data=train_dataset,
     32           validation_freq=1)  # Validate after each epoch
     34 # Save the trained model
     35 model.save(f'chest_xray_model_{dataset_name}.h5')

File /opt/conda/lib/python3.10/site-packages/keras/src/utils/traceback_utils.py:123, in filter_traceback.<locals>.error_handler(*args, **kwargs)
    120     filtered_tb = _process_traceback_frames(e.__traceback__)
    121     # To get the full stack trace, call:
    122     # `keras.config.disable_traceback_filtering()`
---> 123     raise e.with_traceback(filtered_tb) from None
    124 finally:
    125     del filtered_tb

File /opt/conda/lib/python3.10/site-packages/tree/__init__.py:435, in map_structure(func, *structures, **kwargs)
    432 for other in structures[1:]:
    433   assert_same_structure(structures[0], other, check_types=check_types)
    434 return unflatten_as(structures[0],
---> 435                     [func(*args) for args in zip(*map(flatten, structures))])

File /opt/conda/lib/python3.10/site-packages/tree/__init__.py:435, in <listcomp>(.0)
    432 for other in structures[1:]:
    433   assert_same_structure(structures[0], other, check_types=check_types)
    434 return unflatten_as(structures[0],
---> 435                     [func(*args) for args in zip(*map(flatten, structures))])

ValueError: Invalid dtype: object

解决方案

问题根源

错误核心是生成器输出的标签数组 dtype 为 object,而非模型期望的数值类型(如float32或int32)。尽管设置了dtype='int32',但y_col_preprocessor返回的numpy数组可能未被生成器正确转换,同时原代码存在缩进错误导致验证集存储异常。

具体修复步骤

  1. 预编码标签到DataFrame
    放弃依赖生成器的y_col_preprocessor,提前在DataFrame中生成编码后的标签列,确保类型正确:

    # 将无序的set转为有序列表,保证编码一致性
    all_diseases = sorted(list(set('|'.join(df['Finding Labels']).split('|'))))
    num_classes = len(all_diseases)
    
    # 重写标签处理函数,返回float32类型的列表(匹配sigmoid输出)
    def process_labels(label_str):
        labels = label_str.split('|')
        encoded = np.zeros(num_classes, dtype=np.float32)
        for idx, disease in enumerate(all_diseases):
            if disease in labels:
                encoded[idx] = 1.0
        return encoded.tolist()
    
    # 在每个子集DataFrame中添加预编码标签列
    for name, subset_df in datasets.items():
        # ... 保留原有的image_path处理代码 ...
        subset_df['encoded_labels'] = subset_df['Finding Labels'].apply(process_labels)
    
  2. 修改生成器配置
    生成器直接使用预编码的标签列,不再依赖y_col_preprocessor:

    # 训练生成器示例
    train_dataset_generator = train_datagen.flow_from_dataframe(
        dataframe=train_df,
        directory=None,
        target_size=(224, 224),
        batch_size=32,
        class_mode='raw',
        x_col="image_path",
        y_col="encoded_labels",  # 使用预编码列
        shuffle=False,
        interpolation="nearest",
        validate_filenames=True,
        dtype=np.float32  # 与模型输出类型匹配
    )
    
    # 验证生成器同理修改y_col为'encoded_labels'
    
  3. 修复缩进错误
    将validation_datasets[name] = test_dataset_generator缩进至for循环内部,确保每个数据集的验证集都被正确存储:

    for name, subset_df in datasets.items():
        # ... 其他代码 ...
        test_dataset_generator = test_datagen.flow_from_dataframe(...)
        validation_datasets[name] = test_dataset_generator  # 缩进至此行
    
  4. 验证标签类型
    打印生成器的第一个批次标签,确认dtype为数值类型:

    print(train_dataset_generator[0][1].dtype)
    

内容的提问来源于stack exchange,提问作者Nitin Kumar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.27 15:45:56