基于CNN高光谱遥感图像分类项目code如何训练数据集Indianpines基于pytorch框架实现基于pytorch框架实现基于CNN高光谱遥感图像分类项目代码基于PyTorch框架实现一个高光谱遥感图像分类项目并使用Indian Pines数据集。印度Pines数据集_——含166类地物类型和220个波段的光谱信息。文章代码及内容仅供参考环境准备确保您已经安装了以下软件和库Python 3.8 或更高版本PyTorch 1.9 或更高版本torchvision 0.10 或更高版本numpymatplotlibscikit-learnscipy使用以下命令安装所需的Python库pipinstalltorch torchvision numpy matplotlib scikit-learn scipy数据集准备印度Pines数据集可以从官方网站下载。下载后解压文件并将其放置在合适的位置。数据集结构假设数据集解压后的目录结构如下datasets/ └── indian_pines/ ├── Indian_pines_corrected.mat └── Indian_pines_gt.mat数据预处理我们需要加载数据集并进行必要的预处理包括归一化、PCA降维可选、划分训练集和测试集等。加载数据集[titleLoad and Preprocess Indian Pines Dataset]importosimportnumpyasnpimportscipy.ioassiofromsklearn.decompositionimportPCAfromsklearn.model_selectionimporttrain_test_splitfromsklearn.preprocessingimportStandardScalerimporttorchfromtorch.utils.dataimportDataset,DataLoaderfromtorchvisionimporttransformsclassHyperspectralDataset(Dataset):def__init__(self,data,labels,transformNone):self.datadata self.labelslabels self.transformtransformdef__len__(self):returnself.data.shape[0]def__getitem__(self,idx):sampleself.data[idx]labelself.labels[idx]ifself.transform:sampleself.transform(sample)returnsample,labeldefload_indian_pines_data(data_path,gt_path,apply_pcaFalse,n_components50):# Load datadatasio.loadmat(data_path)[indian_pines_corrected]gtsio.loadmat(gt_path)[indian_pines_gt]# Flatten the data and GTdata_flatdata.reshape(-1,data.shape[-1])gt_flatgt.flatten()# Remove unlabeled pixels (label 0)maskgt_flat0data_flatdata_flat[mask]gt_flatgt_flat[mask]-1# Convert to zero-indexed# Normalize datascalerStandardScaler()data_flat_normalizedscaler.fit_transform(data_flat)# Apply PCA if specifiedifapply_pca:pcaPCA(n_componentsn_components)data_flat_reducedpca.fit_transform(data_flat_normalized)else:data_flat_reduceddata_flat_normalized# Reshape back to original shape minus unlabeled pixelsnum_bandsdata_flat_reduced.shape[1]data_shape(np.sum(mask),num_bands)data_reduceddata_flat_reduced.reshape(data_shape)# Split into training and testing setsX_train,X_test,y_train,y_testtrain_test_split(data_reduced,gt_flat,test_size0.2,random_state42,stratifygt_flat)returnX_train,X_test,y_train,y_test# Load datasetdata_path../datasets/indian_pines/Indian_pines_corrected.matgt_path../datasets/indian_pines/Indian_pines_gt.matX_train,X_test,y_train,y_testload_indian_pines_data(data_path,gt_path,apply_pcaTrue,n_components50)# Create datasets and dataloaderstrain_datasetHyperspectralDataset(X_train,y_train,transformtorch.tensor)test_datasetHyperspectralDataset(X_test,y_test,transformtorch.tensor)train_loaderDataLoader(train_dataset,batch_size64,shuffleTrue)test_loaderDataLoader(test_dataset,batch_size64,shuffleFalse)构建CNN模型我们将构建一个简单的卷积神经网络来对高光谱数据进行分类。[titleCNN Model for Hyperspectral Classification]importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassSimpleCNN(nn.Module):def__init__(self,input_dim,hidden_dim,output_dim):super(SimpleCNN,self).__init__()self.fc1nn.Linear(input_dim,hidden_dim)self.fc2nn.Linear(hidden_dim,hidden_dim//2)self.fc3nn.Linear(hidden_dim//2,output_dim)defforward(self,x):xF.relu(self.fc1(x))xF.relu(self.fc2(x))xself.fc3(x)returnx# Define model parametersinput_dimX_train.shape[1]hidden_dim256output_dimlen(np.unique(y_train))modelSimpleCNN(input_dim,hidden_dim,output_dim)训练代码接下来我们将编写训练代码以训练我们的模型。[titleTraining Code for CNN on Hyperspectral Data]importtorch.optimasoptimfromsklearn.metricsimportclassification_report,accuracy_score# Define loss function and optimizercriterionnn.CrossEntropyLoss()optimizeroptim.Adam(model.parameters(),lr0.001)# Training loopnum_epochs100devicetorch.device(cudaiftorch.cuda.is_available()elsecpu)model.to(device)forepochinrange(num_epochs):model.train()running_loss0.0forinputs,labelsintrain_loader:inputsinputs.float().to(device)labelslabels.long().to(device)optimizer.zero_grad()outputsmodel(inputs)losscriterion(outputs,labels)loss.backward()optimizer.step()running_lossloss.item()avg_lossrunning_loss/len(train_loader)print(fEpoch [{epoch1}/{num_epochs}], Train Loss:{avg_loss:.4f})# Evaluation on validation setmodel.eval()correct0total0all_preds[]all_labels[]withtorch.no_grad():forinputs,labelsintest_loader:inputsinputs.float().to(device)labelslabels.long().to(device)outputsmodel(inputs)_,predictedtorch.max(outputs.data,1)totallabels.size(0)correct(predictedlabels).sum().item()all_preds.extend(predicted.cpu().numpy())all_labels.extend(labels.cpu().numpy())accuracycorrect/totalprint(fEpoch [{epoch1}/{num_epochs}], Test Accuracy:{accuracy:.4f})# Print classification reportprint(classification_report(all_labels,all_preds,target_names[str(i)foriinrange(output_dim)]))模型评估在训练过程中计算了准确率和其他指标。为了更好地可视化模型性能我们可以绘制混淆矩阵。[titleConfusion Matrix Visualization]fromsklearn.metricsimportconfusion_matriximportseabornassnsimportmatplotlib.pyplotasplt# Compute confusion matrixconf_matconfusion_matrix(all_labels,all_preds)# Plot confusion matrixplt.figure(figsize(20,16))sns.heatmap(conf_mat,annotTrue,fmtd,cmapBlues,xticklabels[str(i)foriinrange(output_dim)],yticklabels[str(i)foriinrange(output_dim)])plt.xlabel(Predicted)plt.ylabel(True)plt.title(Confusion Matrix)plt.show()总结通过上述步骤我们可以构建一个基于PyTorch的高光谱遥感图像分类系统使用印度Pines数据集进行训练和评估。以下是所有相关的代码文件数据加载和预处理(load_and_preprocess.py)CNN模型实现(simple_cnn.py)训练代码(training_code.py)混淆矩阵可视化(confusion_matrix.py)