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

在Keras ResNet50模型中获取每批次训练的图像名称

解决Keras微调ResNet50时记录每批次训练图像名称的问题

这个问题我之前微调模型时也碰到过,确实Keras原生的fit或fit_generator没法直接传递图像名称这类额外信息,但我们可以通过自定义数据生成器+回调函数或者自定义训练循环这两种方式来实现需求,下面给你详细拆解:

方法一:自定义Sequence生成器 + 回调函数

这种方法最贴合Keras的原生训练流程,不需要大幅修改训练代码,核心思路是让生成器在输出图像和标签的同时,保存当前批次的文件名,再用回调函数在每批次训练结束后把文件名写入日志。

步骤1:实现自定义图像生成器

继承keras.utils.Sequence类,在生成批次数据时记录当前批次的文件名:

from tensorflow.keras.utils import Sequence
import numpy as np
import cv2

class CustomImageGenerator(Sequence):
    def __init__(self, image_paths, labels, batch_size, img_size=(224,224)):
        self.image_paths = image_paths  # 存储所有图像的完整路径
        self.labels = labels
        self.batch_size = batch_size
        self.img_size = img_size
        self.current_batch_filenames = []  # 用于临时存储当前批次的文件名

    def __len__(self):
        # 返回总批次数
        return int(np.ceil(len(self.image_paths) / self.batch_size))

    def __getitem__(self, idx):
        # 截取当前批次的图像路径和标签
        batch_paths = self.image_paths[idx*self.batch_size : (idx+1)*self.batch_size]
        batch_labels = self.labels[idx*self.batch_size : (idx+1)*self.batch_size]
        
        # 提取文件名(这里假设路径是类似"dataset/cat/1.jpg",截取最后一部分)
        self.current_batch_filenames = [path.split('/')[-1] for path in batch_paths]
        
        # 加载并预处理图像(也可以用ResNet50自带的preprocess_input)
        batch_images = []
        for path in batch_paths:
            img = cv2.imread(path)
            img = cv2.resize(img, self.img_size)
            img = img / 255.0  # 简单归一化,也可以替换成ResNet的预处理
            batch_images.append(img)
        
        return np.array(batch_images), np.array(batch_labels)

步骤2:实现批次日志回调函数

自定义回调函数,在每批次训练结束后读取生成器中保存的文件名,写入日志文件:

from tensorflow.keras.callbacks import Callback

class BatchImageLogger(Callback):
    def __init__(self, generator, log_file='batch_training_logs.log'):
        self.generator = generator
        self.log_file = log_file
        # 初始化日志文件,写入表头
        with open(self.log_file, 'w', encoding='utf-8') as f:
            f.write('Epoch,Batch Index,Image Filenames\n')

    def on_train_batch_end(self, batch, logs=None):
        # 获取当前训练的epoch数
        current_epoch = self.model.optimizer.iterations.numpy() // len(self.generator)
        # 获取当前批次的文件名
        filenames = self.generator.current_batch_filenames
        # 写入日志,用逗号分隔多个文件名
        with open(self.log_file, 'a', encoding='utf-8') as f:
            f.write(f'{current_epoch},{batch},{",".join(filenames)}\n')

步骤3:整合到训练流程

用你的ResNet50模型配合自定义生成器和回调函数训练:

# 假设你已经准备好image_paths(所有图像路径列表)和labels(对应标签列表)
train_generator = CustomImageGenerator(image_paths, labels, batch_size=32)
batch_logger = BatchImageLogger(train_generator)

# 构建微调的ResNet50模型
from tensorflow.keras.applications import ResNet50
from tensorflow.keras.layers import Dense, GlobalAveragePooling2D
from tensorflow.keras.models import Model

base_model = ResNet50(weights='imagenet', include_top=False)
x = base_model.output
x = GlobalAveragePooling2D()(x)
x = Dense(1024, activation='relu')(x)
predictions = Dense(num_classes, activation='softmax')(x)  # num_classes是你的分类数
model = Model(inputs=base_model.input, outputs=predictions)

# 冻结预训练层,先训练顶层
for layer in base_model.layers:
    layer.trainable = False

model.compile(optimizer='adam', loss='categorical_crossentropy', metrics=['accuracy'])

# 启动训练,传入回调函数
model.fit(train_generator, epochs=10, callbacks=[batch_logger])

方法二:自定义训练循环(更灵活)

如果需要对训练过程有更强的控制权,可以用TensorFlow的自定义训练循环,直接在每一步获取批次的文件名并记录,不需要依赖回调函数。

示例代码

import tensorflow as tf

# 定义图像加载和预处理函数,返回图像、标签、文件名
def load_preprocess_image(file_path, label):
    img = tf.io.read_file(file_path)
    img = tf.image.decode_jpeg(img, channels=3)
    img = tf.image.resize(img, (224, 224))
    img = tf.keras.applications.resnet50.preprocess_input(img)
    return img, label, file_path

# 构建tf.data.Dataset
image_paths = ["path/to/img1.jpg", "path/to/img2.jpg", ...]  # 你的图像路径列表
labels = [0, 1, ...]  # 对应标签列表
dataset = tf.data.Dataset.from_tensor_slices((image_paths, labels))
dataset = dataset.map(load_preprocess_image, num_parallel_calls=tf.data.AUTOTUNE)
dataset = dataset.batch(32).prefetch(tf.data.AUTOTUNE)

# 构建ResNet50微调模型(和方法一一致)
base_model = ResNet50(weights='imagenet', include_top=False)
x = base_model.output
x = GlobalAveragePooling2D()(x)
x = Dense(1024, activation='relu')(x)
predictions = Dense(num_classes, activation='softmax')(x)
model = Model(inputs=base_model.input, outputs=predictions)

for layer in base_model.layers:
    layer.trainable = False

# 定义优化器、损失函数和指标
optimizer = tf.keras.optimizers.Adam()
loss_fn = tf.keras.losses.CategoricalCrossentropy()
accuracy_metric = tf.keras.metrics.CategoricalAccuracy()

# 初始化日志文件
log_file = 'batch_training_logs.log'
with open(log_file, 'w', encoding='utf-8') as f:
    f.write('Epoch,Batch Index,Image Filenames\n')

# 自定义训练循环
epochs = 10
for epoch in range(epochs):
    print(f"=== Epoch {epoch+1}/{epochs} ===")
    batch_idx = 0
    for batch_data in dataset:
        imgs, labels, filenames = batch_data
        
        # 将文件名张量转换为字符串列表
        filenames = [fname.numpy().decode('utf-8').split('/')[-1] for fname in filenames]
        
        # 前向传播+反向传播
        with tf.GradientTape() as tape:
            preds = model(imgs, training=True)
            loss = loss_fn(labels, preds)
        
        gradients = tape.gradient(loss, model.trainable_variables)
        optimizer.apply_gradients(zip(gradients, model.trainable_variables))
        
        # 更新准确率指标
        accuracy_metric.update_state(labels, preds)
        
        # 写入日志
        with open(log_file, 'a', encoding='utf-8') as f:
            f.write(f'{epoch},{batch_idx},{",".join(filenames)}\n')
        
        batch_idx += 1
    
    # 打印当前epoch的准确率
    print(f"Epoch {epoch+1} Accuracy: {accuracy_metric.result().numpy():.4f}")
    accuracy_metric.reset_states()

两种方法的选择

  • 如果习惯用Keras原生的fit方法,优先选方法一,代码改动小,符合常规训练流程;
  • 如果需要在训练过程中做更多自定义操作(比如根据文件名做特殊处理),选方法二,灵活性更高。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.25 07:16:01