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

ML.NET图像分类训练Accuracy为NaN且预测值一致问题排查

问题

我尝试加载FashionMNIST数据集训练图像分类模型,模型相关类定义如下:

public class DataImage
{
    [LoadColumn(0)]
    [ColumnName("Image")]
    public byte[] Image { get; set; }

    [LoadColumn(1)]
    [ColumnName("Label")]
    public int Label { get; set; }

}

public class DataImagePredict
{
    [LoadColumn(1)]
    [ColumnName("Label")]
    public int Label { get; set; }

    [LoadColumn(2)]
    [ColumnName(@"PredictedLabel")]
    public int PredictedLabel { get; set; }

    [LoadColumn(3)]
    [ColumnName(@"Score")]
    public float[] Score { get; set; }

    public override string ToString()
    {
        return $"PredictedLabel : {PredictedLabel} Score {string.Join(",", Score)} ";
    }
}

训练代码如下:

DataImage[] trainingData = FashionMNISTML.LoadTrainImages();

DataImage[] testData = FashionMNISTML.LoadTestImages();

var mlContext = new MLContext();

IDataView trainingDataView = mlContext.Data.LoadFromEnumerable(trainingData);

IDataView testDataView = mlContext.Data.LoadFromEnumerable(testData);

trainingDataView = mlContext.Transforms.Conversion
            .MapValueToKey("Label", "Label", keyOrdinality: ValueToKeyMappingEstimator.KeyOrdinality.ByValue)
    .Fit(trainingDataView)
    .Transform(trainingDataView);

testDataView = mlContext.Transforms.Conversion
            .MapValueToKey("Label", "Label", keyOrdinality: ValueToKeyMappingEstimator.KeyOrdinality.ByValue).Fit(testDataView).Transform(testDataView);

var options = new ImageClassificationTrainer.Options()
{
    FeatureColumnName = "Image",
    LabelColumnName = "Label",
    Arch = ImageClassificationTrainer.Architecture.ResnetV250,
    Epoch = 10,
    BatchSize = 10,
    LearningRate = 0.01f,
    MetricsCallback = (metrics) => Console.WriteLine(metrics),
    EarlyStoppingCriteria = null
};

var trainerPipLine =   mlContext.MulticlassClassification.Trainers.ImageClassification(options)
                        .Append(mlContext.Transforms.Conversion.MapKeyToValue(
                                outputColumnName : "PredictedLabel",
                                inputColumnName : "PredictedLabel"
                            )) ;


ITransformer model = trainerPipLine.Fit(trainingDataView);

其中FashionMNIST.LoadTrainingImages()用于读取未压缩文件,训练集含60000张图像,测试集含10000张。训练过程中输出如下:

Phase: Training, Dataset used:      Train, Batch Processed Count:   0, Epoch:   0, Accuracy:        NaN, Cross-Entropy:        NaN, Learning Rate:       0.01
Phase: Training, Dataset used:      Train, Batch Processed Count:   0, Epoch:   9, Accuracy:        NaN, Cross-Entropy:        NaN, Learning Rate:       0.01

所有预测结果一致且Score为0。当前使用最新版ML.NET及SciSharp.TensorFlow 2.3.1(因此前问题需使用该版本),请问问题出在哪里?

分析与解决

1. 图像数据格式不兼容

ImageClassificationTrainer无法直接识别byte[]类型的图像数据,它需要ML.NET规范的结构化图像格式(包含高度、宽度、通道信息)。直接传入原始字节数组会导致模型无法解析图像张量,进而无法计算有效损失,出现NaN指标和无效预测。

解决方法:
添加图像格式转换步骤,将byte[]转为ML.NET可识别的图像结构,同时指定FashionMNIST的图像参数(28x28单通道灰度图):

// 在Label映射前添加图像加载转换
var imageLoadTransformer = mlContext.Transforms.LoadRawImageBytes(
        outputColumnName: "Image",
        imageFolder: "", // 因读取的是字节流,此处留空
        inputColumnName: nameof(DataImage.Image),
        height: 28,
        width: 28,
        channels: 1
    ).Fit(trainingDataView);

trainingDataView = imageLoadTransformer.Transform(trainingDataView);
testDataView = imageLoadTransformer.Transform(testDataView);

2. Label映射规则不一致

分别对训练集和测试集单独拟合MapValueToKey转换器,会导致两者的标签-键映射规则可能不统一(比如训练集标签0对应键0,测试集可能因数据分布差异映射到其他值),直接影响模型预测的正确性。

解决方法:
复用训练集的Label映射转换器处理测试集:

// 仅在训练集上拟合映射转换器
var labelMapTransformer = mlContext.Transforms.Conversion
    .MapValueToKey("Label", "Label", keyOrdinality: ValueToKeyMappingEstimator.KeyOrdinality.ByValue)
    .Fit(trainingDataView);

// 用同一个转换器处理训练集和测试集
trainingDataView = labelMapTransformer.Transform(trainingDataView);
testDataView = labelMapTransformer.Transform(testDataView);

3. 图像尺寸不匹配预训练模型要求

你使用的ResnetV250预训练模型要求输入图像尺寸为224x224,但FashionMNIST是28x28的小图,直接输入会导致张量维度不匹配,模型无法正常计算,进而出现训练无进展的情况。

解决方法:
添加图像缩放步骤,将小图放大到预训练模型要求的尺寸:

// 在图像加载后添加resize转换
var resizeTransformer = mlContext.Transforms.ResizeImages(
        outputColumnName: "Image",
        imageWidth: 224,
        imageHeight: 224,
        inputColumnName: "Image"
    ).Fit(trainingDataView);

trainingDataView = resizeTransformer.Transform(trainingDataView);
testDataView = resizeTransformer.Transform(testDataView);

4. 训练数据未被正常迭代

训练日志中Batch Processed Count始终为0,本质是前面的数据格式错误导致训练器无法读取和迭代数据,解决上述图像格式、尺寸问题后,该指标会正常更新,模型也能进入有效训练阶段。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 04:42:05