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

如何在CNN全量训练集上运行PCA并将降维结果输入CNN

背景

我有一个分类准确率为98%的CNN模型,训练时长约为2分钟。我希望在训练CNN之前,对训练集执行PCA操作来减少训练时长,目标是将训练时间降低到1分钟甚至更短。

问题

目前我遇到的核心问题是:我不知道如何对我的3万张训练图像全量运行PCA,再将处理后的图像传入CNN。

  • 我已经在训练集的几百张样本图像上成功运行了PCA,但不清楚如何扩展到全量训练集上执行。
  • 除此之外,即便我完成了所有训练图像的PCA处理,我要如何将PCA的输出和CNN的输入“连接”起来?换句话说,我要如何将PCA输出的低维重建图像喂入CNN?

我已经在Stack Overflow和全网大范围搜索相关示例或同类问题,但没有找到可用的解决方案,如果有人能提供帮助我将非常感激。

以下是我的数据集中部分样本图像经过PCA处理后的效果:
PCA处理后样本图像

最小可复现示例
pip install tensorflow
pip install numpy
pip install matplotlib

"""# Import Libraries"""

# Import Libraries
import tensorflow as tf
from tensorflow import keras
from keras.models import Sequential
from keras.layers import Dense, Flatten, Conv2D, MaxPooling2D, Dropout
from tensorflow.keras import layers
from tensorflow.keras.utils import to_categorical
import numpy as np
import matplotlib.pyplot as plt

plt.style.use('fivethirtyeight')

"""# Load Dataset"""

import pathlib
dataset_url = "*/TrainingSet.tar.gz"
data_dir = tf.keras.utils.get_file(origin = dataset_url,
                                   fname = "TrainingSet",
                                   untar = True)
data_dir = pathlib.Path(data_dir)

"""# Display # Images to check"""

print(list(data_dir.glob('*/*.png')))
image_count = len(list(data_dir.glob('*/*.png')))
print(image_count)

"""# Display sample image"""

pip install sklearn

import numpy as np
import os
import PIL
import PIL.Image
import tensorflow as tf
import tensorflow_datasets as tfds
from sklearn.decomposition import PCA

graphs = list(data_dir.glob('*/*.png'))
PIL.Image.open(str(graphs[6]))

"""# Define Image Dimensions & Batch Size"""

batch_size = 32
img_height = 36
img_width = 36

"""# Create Training & Validation Sets (80%, 20%)"""

train_ds = tf.keras.preprocessing.image_dataset_from_directory(
  data_dir,
  validation_split=0.2,
  subset="training",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)

val_ds = tf.keras.preprocessing.image_dataset_from_directory(
  data_dir,
  validation_split=0.2,
  subset="validation",
  seed=123,
  image_size=(img_height, img_width),
  batch_size=batch_size)

"""# Define 3 Classes"""

class_names = ['Cubic Sinusoidal', 'Linear Sinusoidal', 'Quadratic Sinusoidal']
print(class_names)

"""# Supervised Learning (9 Samples from the Training Set)"""

!pip install skimage

from skimage import data
from skimage.color import rgb2gray

import matplotlib.pyplot as plt

subGraphs = []

plt.figure(figsize=(10, 10))
for images, labels in train_ds.take(1):
  for i in range(9):
    ax = plt.subplot(3, 3, i + 1)
    plt.imshow(images[i].numpy().astype("uint8"))
    subGraphs.append(images[i].numpy().astype("uint8"))
    plt.title(class_names[labels[i]])
    plt.axis("off")

subGraphs = np.array(subGraphs)
print(subGraphs.shape)

grayscale = rgb2gray(subGraphs[1])
print(grayscale.shape)

X=grayscale 

pca_oliv = PCA(n_components = 36)
X_proj = pca_oliv.fit_transform(X)

print(np.cumsum(pca_oliv.explained_variance_ratio_))
plt.plot(np.cumsum(pca_oliv.explained_variance_ratio_))

plt.imshow(np.reshape(pca_oliv.components_, (36,36)), cmap=plt.cm.bone, interpolation='nearest')

X_inv_proj = pca_oliv.inverse_transform(X_proj)
X_proj_img = np.reshape(X_inv_proj,(1,36,36))

plt.imshow(X_proj_img[0], cmap=plt.cm.bone, interpolation='nearest')

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 08:48:03