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

trainNetwork报错:X与Y样本数量不匹配,请求排查解决

MATLAB trainNetwork报错:X与Y样本数不匹配的排查与解决

问题描述

运行MATLAB代码训练网络时触发报错:

Error using trainNetwork Number of observations in X and Y disagree

通过以下代码检查尺寸:

disp(['Size of featuresArray: ' num2str(size(featuresArray))]);
disp(['Size of trainingLabels: ' num2str(size(trainingLabels))]);

得到结果:

Size of featuresArray: 957  100
Size of trainingLabels: 957    1

但训练仍失败,原始完整代码如下:

% Import required libraries
import matlab.io.datastore.ImageDatastore
import matlab.io.datastore.*
import matlab.io.datastore.augmenters.image.*
import matlab.io.datastore.augmenters.*
import matlab.io.datastore.TransformedDatastore.*

% Set the folder containing the image dataset
datasetFolder = 'C:\Users\xinai\Downloads\Images';

% Create an imageDatastore for the dataset
imds = imageDatastore(datasetFolder, 'IncludeSubfolders', true, 'LabelSource', 'foldernames');

% Split the dataset into training and validation sets
[trainingData, validationData] = splitEachLabel(imds, 0.7, 'randomized');

% Define the rotation range for augmentation
rotationRange = [-10 10];

% Define the input size
inputSize = [224 224 3];

% Define the number of GLCM features
numGLCMFeatures = 5;

% Define the number of HSV histogram bins
numHSVHistBins = 32;

% Define the number of PCA components
numComponents = 100;

% Define the preprocessing function
preprocessFcn = @(img) {im2gray(img), hsvHistogram(img, numHSVHistBins)};

% Initialize feature matrices
numObservations = numel(trainingData.Files);
glcmFeatures = zeros(numObservations, numGLCMFeatures);
hsvHistograms = zeros(numObservations, numHSVHistBins * 3);

% Extract GLCM features and HSV histograms
reset(trainingData);
for i = 1:numObservations
    % Read the image
    img = read(trainingData);

    % Apply data augmentation
    imgAugmented = augmentedImageDatastore(inputSize, img, ...
        'DataAugmentation', imageDataAugmenter('RandRotation', rotationRange, ...
        'RandXReflection', true, 'RandYReflection', true, ...
        'RandXScale', [0.8 1.2], 'RandYScale', [0.8 1.2]));

     % Read the image
    img = readimage(trainingData, i);

    % Apply data augmentation
    augmentedData = augment(augmenter, img);

    % Convert the image to grayscale
    grayImg = rgb2gray(img);

    % Calculate GLCM features
    glcm = graycomatrix(grayImg);
    stats = graycoprops(glcm);
    glcmFeatures(i, :) = [stats.Contrast, stats.Correlation, stats.Energy, stats.Homogeneity, stats.Homogeneity];

    % Calculate HSV histogram
    hsvImg = rgb2hsv(img);
    hHist = imhist(hsvImg(:, :, 1), numHSVHistBins);
    sHist = imhist(hsvImg(:, :, 2), numHSVHistBins);
    vHist = imhist(hsvImg(:, :, 3), numHSVHistBins);
    hsvHistograms(i, :) = [hHist', sHist', vHist'];
end

% Concatenate GLCM features and HSV histograms
features = [glcmFeatures, hsvHistograms];

% Perform PCA on the features
[coeff, score, ~, ~, explained] = pca(features);
selectedComponents = coeff(:, 1:numComponents);
projectedFeatures = score(:, 1:numComponents);

% Create a new table with the projected features and labels
pcaTrainingData = table(projectedFeatures, trainingData.Labels, 'VariableNames', {'Features', 'Labels'});

% Select a subset of training labels to match the number of observations in projectedFeatures
numObservations = size(pcaTrainingData, 1);
pcaTrainingData = pcaTrainingData(1:numObservations, :);

% Convert the training labels to categorical
trainingLabels = categorical(pcaTrainingData.Labels);

% Convert the features table to an array
featuresArrayTranspose = projectedFeatures;

% Check the sizes of featuresArrayTranspose and trainingLabels
size(featuresArrayTranspose)
size(trainingLabels)


% Load the ResNet-50 network
net = resnet50;

% Replace the fully connected layer with a new layer
numClasses = numel(categories(trainingLabels));
newFullyConnectedLayer = fullyConnectedLayer(numClasses, 'Name', 'fc', 'WeightLearnRateFactor', 10, 'BiasLearnRateFactor', 10);

% Remove the last layers
lgraph = layerGraph(net);
lgraph = removeLayers(lgraph, {'fc1000', 'fc1000_softmax', 'ClassificationLayer_fc1000'});

% Connect the new fully connected layer to the network
newFullyConnectedLayer = fullyConnectedLayer(numClasses, 'Name', 'fc', 'WeightLearnRateFactor', 10, 'BiasLearnRateFactor', 10);
lgraph = addLayers(lgraph, newFullyConnectedLayer);
lgraph = connectLayers(lgraph, 'avg_pool', 'fc');

% Set the output layer
newSoftmaxLayer = softmaxLayer('Name', 'softmax');
newClassificationLayer = classificationLayer('Name', 'ClassificationLayer');
lgraph = addLayers(lgraph, newSoftmaxLayer);
lgraph = addLayers(lgraph, newClassificationLayer);
lgraph = connectLayers(lgraph, 'fc', 'softmax');
lgraph = connectLayers(lgraph, 'softmax', 'ClassificationLayer');

% Set training options
miniBatchSize = 10;
numIterationsPerEpoch = floor(numel(trainingData.Files) / miniBatchSize);
trainOpts = trainingOptions('sgdm', ...
    'MiniBatchSize', miniBatchSize, ...
    'MaxEpochs', 50, ...
    'InitialLearnRate', 1e-3, ...
    'Shuffle', 'every-epoch', ...
    'ValidationData', validationData, ...
    'ValidationFrequency', numIterationsPerEpoch, ...
    'Verbose', false, ...
    'Plots', 'training-progress');

% Select a subset of featuresArray to match the number of observations in trainingLabels
featuresArraySubset = featuresArrayTranspose(1:numObservations, :);

% Transpose the features array to match the network's input requirements
featuresArrayTranspose = featuresArraySubset';

% Check the sizes of featuresArrayTranspose and trainingLabels
size(featuresArrayTranspose)
size(trainingLabels)

% Convert the training labels to categorical
trainingLabels = categorical(trainingLabels);

% Train the network
net = trainNetwork(featuresArrayTranspose, trainingLabels, lgraph, trainOpts);

关键错误点分析

  • 特征矩阵维度错误:trainNetwork针对特征向量输入时,要求输入矩阵格式为样本数×特征数(即N×F),但代码中错误地将特征矩阵转置为100×957,导致网络识别到的样本数(100)与标签数(957)不匹配,触发报错。
  • 数据增强逻辑失效:代码中重复读取图像,且未定义augmenter就调用augment(augmenter, img),增强后的图像也未用于特征提取,完全冗余。
  • 验证数据格式不匹配:训练用提取的特征矩阵,而验证数据使用原始ImageDatastore,两者数据格式不一致,trainNetwork无法处理混合格式的训练/验证数据。
  • 冗余的标签子集选择:pcaTrainingData = pcaTrainingData(1:numObservations, :);属于无效代码,pcaTrainingData本身就是样本数对应的标签,无需额外截取。

修正步骤

  1. 保持特征矩阵为样本数×特征数的格式,删除错误的转置操作。
  2. 修复数据增强流程:定义augmenter,使用增强后的图像提取特征。
  3. 对验证集执行与训练集完全一致的特征提取+PCA投影,确保验证数据格式与训练集匹配。
  4. 删除冗余的无效代码。

完整修正代码

% Import required libraries
import matlab.io.datastore.ImageDatastore
import matlab.io.datastore.*
import matlab.io.datastore.augmenters.image.*
import matlab.io.datastore.augmenters.*

% Set the folder containing the image dataset
datasetFolder = 'C:\Users\xinai\Downloads\Images';

% Create an imageDatastore for the dataset
imds = imageDatastore(datasetFolder, 'IncludeSubfolders', true, 'LabelSource', 'foldernames');

% Split the dataset into training and validation sets
[trainingData, validationData] = splitEachLabel(imds, 0.7, 'randomized');

% Define the rotation range for augmentation
rotationRange = [-10 10];
% 定义数据增强器
augmenter = imageDataAugmenter('RandRotation', rotationRange, ...
    'RandXReflection', true, 'RandYReflection', true, ...
    'RandXScale', [0.8 1.2], 'RandYScale', [0.8 1.2]);

% Define the input size
inputSize = [224 224 3];

% Define the number of GLCM features
numGLCMFeatures = 5;

% Define the number of HSV histogram bins
numHSVHistBins = 32;

% Define the number of PCA components
numComponents = 100;

% 定义特征提取函数
function [glcmFeat, hsvFeat] = extractFeatures(img, numGLCMFeatures, numHSVHistBins)
    % Convert the image to grayscale
    grayImg = rgb2gray(img);
    % Calculate GLCM features
    glcm = graycomatrix(grayImg);
    stats = graycoprops(glcm);
    glcmFeat = [stats.Contrast, stats.Correlation, stats.Energy, stats.Homogeneity, stats.Homogeneity];
    % Calculate HSV histogram
    hsvImg = rgb2hsv(img);
    hHist = imhist(hsvImg(:, :, 1), numHSVHistBins);
    sHist = imhist(hsvImg(:, :, 2), numHSVHistBins);
    vHist = imhist(hsvImg(:, :, 3), numHSVHistBins);
    hsvFeat = [hHist', sHist', vHist'];
end

% ---------------------- 处理训练集 ----------------------
% Initialize feature matrices
numTrainObs = numel(trainingData.Files);
glcmFeatures_train = zeros(numTrainObs, numGLCMFeatures);
hsvHistograms_train = zeros(numTrainObs, numHSVHistBins * 3);

% Extract GLCM features and HSV histograms
for i = 1:numTrainObs
    % Read the image
    img = readimage(trainingData, i);
    % Apply data augmentation
    imgAugmented = augment(augmenter, img);
    % Extract features using augmented image
    [glcmFeat, hsvFeat] = extractFeatures(imgAugmented, numGLCMFeatures, numHSVHistBins);
    glcmFeatures_train(i, :) = glcmFeat;
    hsvHistograms_train(i, :) = hsvFeat;
end

% Concatenate GLCM features and HSV histograms
features_train = [glcmFeatures_train, hsvHistograms_train];

% Perform PCA on the training features
[coeff, score_train, ~, ~, explained] = pca(features_train);
selectedComponents = coeff(:, 1:numComponents);
projectedFeatures_train = score_train(:, 1:numComponents);

% Prepare training labels
trainingLabels = categorical(trainingData.Labels);

% ---------------------- 处理验证集 ----------------------
% Initialize feature matrices for validation
numValObs = numel(validationData.Files);
glcmFeatures_val = zeros(numValObs, numGLCMFeatures);
hsvHistograms_val = zeros(numValObs, numHSVHistBins * 3);

% Extract features for validation set
for i = 1:numValObs
    img = readimage(validationData, i);
    % 验证集无需数据增强(或按需添加)
    [glcmFeat, hsvFeat] = extractFeatures(img, numGLCMFeatures, numHSVHistBins);
    glcmFeatures_val(i, :) = glcmFeat;
    hsvHistograms_val(i, :) = hsvFeat;
end

% Concatenate features and project using training PCA coeff
features_val = [glcmFeatures_val, hsvHistograms_val];
projectedFeatures_val = (features_val - mean(features_train,1)) * selectedComponents;

% Prepare validation labels
validationLabels = categorical(validationData.Labels);
validationData_processed = {projectedFeatures_val, validationLabels};

% ---------------------- 构建网络 ----------------------
% Load the ResNet-50 network
net = resnet50;

% Replace the fully connected layer with a new layer
numClasses = numel(categories(trainingLabels));

% Remove the last layers
lgraph = layerGraph(net);
lgraph = removeLayers(lgraph, {'fc1000', 'fc1000_softmax', 'ClassificationLayer_fc1000'});

% Add new layers
newFullyConnectedLayer = fullyConnectedLayer(numClasses, 'Name', 'fc', 'WeightLearnRateFactor', 10, 'BiasLearnRateFactor', 10);
newSoftmaxLayer = softmaxLayer('Name', 'softmax');
newClassificationLayer = classificationLayer('Name', 'ClassificationLayer');

lgraph = addLayers(lgraph, [newFullyConnectedLayer, newSoftmaxLayer, newClassificationLayer]);
% Connect layers
lgraph = connectLayers(lgraph, 'avg_pool', 'fc');
lgraph = connectLayers(lgraph, 'fc', 'softmax');
lgraph = connectLayers(lgraph, 'softmax', 'ClassificationLayer');

% ---------------------- 设置训练选项并训练 ----------------------
miniBatchSize = 10;
numIterationsPerEpoch = floor(numTrainObs / miniBatchSize);
trainOpts = trainingOptions('sgdm', ...
    'MiniBatchSize', miniBatchSize, ...
    'MaxEpochs', 50, ...
    'InitialLearnRate', 1e-3, ...
    'Shuffle', 'every-epoch', ...
    'ValidationData', validationData_processed, ...
    'ValidationFrequency', numIterationsPerEpoch, ...
    'Verbose', true, ...
    'Plots', 'training-progress');

% 检查维度(可选)
disp(['Size of training features: ' num2str(size(projectedFeatures_train))]);
disp(['Size of training labels: ' num2str(size(trainingLabels))]);

% Train the network
net = trainNetwork(projectedFeatures_train, trainingLabels, lgraph, trainOpts);

额外说明

  • 训练集使用数据增强提升泛化能力,验证集通常不使用增强以准确评估模型性能。
  • PCA投影时,验证集必须使用训练集计算得到的均值和系数,避免数据泄露。
  • 修正后特征矩阵维度为957×100,标签维度为957×1,两者样本数完全匹配,可正常训练。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 00:44:50