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本身就是样本数对应的标签,无需额外截取。
修正步骤
- 保持特征矩阵为样本数×特征数的格式,删除错误的转置操作。
- 修复数据增强流程:定义
augmenter,使用增强后的图像提取特征。 - 对验证集执行与训练集完全一致的特征提取+PCA投影,确保验证数据格式与训练集匹配。
- 删除冗余的无效代码。
完整修正代码
% 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
相关产品推荐
相关产品推荐

