于CNN高光谱遥感图像分类项目code,如何训练数据集,Indianpines,基于pytorch框架实现
2026/7/21 17:39:12 网站建设 项目流程

基于CNN高光谱遥感图像分类项目code,如何训练数据集,Indianpines,基于pytorch框架实现
基于pytorch框架实现基于CNN高光谱遥感图像分类项目代码

基于PyTorch框架实现一个高光谱遥感图像分类项目,并使用Indian Pines数据集。印度Pines数据集_——含166类地物类型和220个波段的光谱信息。

文章代码及内容仅供参考!

环境准备

确保您已经安装了以下软件和库:

  • Python 3.8 或更高版本
  • PyTorch 1.9 或更高版本
  • torchvision 0.10 或更高版本
  • numpy
  • matplotlib
  • scikit-learn
  • scipy

使用以下命令安装所需的Python库:

pipinstalltorch torchvision numpy matplotlib scikit-learn scipy

数据集准备

印度Pines数据集可以从官方网站下载。下载后,解压文件并将其放置在合适的位置。

数据集结构

假设数据集解压后的目录结构如下:

datasets/ └── indian_pines/ ├── Indian_pines_corrected.mat └── Indian_pines_gt.mat

数据预处理

我们需要加载数据集并进行必要的预处理,包括归一化、PCA降维(可选)、划分训练集和测试集等。

加载数据集
[<title="Load and Preprocess Indian Pines Dataset">]importosimportnumpyasnpimportscipy.ioassiofromsklearn.decompositionimportPCAfromsklearn.model_selectionimporttrain_test_splitfromsklearn.preprocessingimportStandardScalerimporttorchfromtorch.utils.dataimportDataset,DataLoaderfromtorchvisionimporttransformsclassHyperspectralDataset(Dataset):def__init__(self,data,labels,transform=None):self.data=data self.labels=labels self.transform=transformdef__len__(self):returnself.data.shape[0]def__getitem__(self,idx):sample=self.data[idx]label=self.labels[idx]ifself.transform:sample=self.transform(sample)returnsample,labeldefload_indian_pines_data(data_path,gt_path,apply_pca=False,n_components=50):# Load datadata=sio.loadmat(data_path)['indian_pines_corrected']gt=sio.loadmat(gt_path)['indian_pines_gt']# Flatten the data and GTdata_flat=data.reshape(-1,data.shape[-1])gt_flat=gt.flatten()# Remove unlabeled pixels (label == 0)mask=gt_flat>0data_flat=data_flat[mask]gt_flat=gt_flat[mask]-1# Convert to zero-indexed# Normalize datascaler=StandardScaler()data_flat_normalized=scaler.fit_transform(data_flat)# Apply PCA if specifiedifapply_pca:pca=PCA(n_components=n_components)data_flat_reduced=pca.fit_transform(data_flat_normalized)else:data_flat_reduced=data_flat_normalized# Reshape back to original shape minus unlabeled pixelsnum_bands=data_flat_reduced.shape[1]data_shape=(np.sum(mask),num_bands)data_reduced=data_flat_reduced.reshape(data_shape)# Split into training and testing setsX_train,X_test,y_train,y_test=train_test_split(data_reduced,gt_flat,test_size=0.2,random_state=42,stratify=gt_flat)returnX_train,X_test,y_train,y_test# Load datasetdata_path='../datasets/indian_pines/Indian_pines_corrected.mat'gt_path='../datasets/indian_pines/Indian_pines_gt.mat'X_train,X_test,y_train,y_test=load_indian_pines_data(data_path,gt_path,apply_pca=True,n_components=50)# Create datasets and dataloaderstrain_dataset=HyperspectralDataset(X_train,y_train,transform=torch.tensor)test_dataset=HyperspectralDataset(X_test,y_test,transform=torch.tensor)train_loader=DataLoader(train_dataset,batch_size=64,shuffle=True)test_loader=DataLoader(test_dataset,batch_size=64,shuffle=False)

构建CNN模型

我们将构建一个简单的卷积神经网络来对高光谱数据进行分类。

[<title="CNN Model for Hyperspectral Classification">]importtorchimporttorch.nnasnnimporttorch.nn.functionalasFclassSimpleCNN(nn.Module):def__init__(self,input_dim,hidden_dim,output_dim):super(SimpleCNN,self).__init__()self.fc1=nn.Linear(input_dim,hidden_dim)self.fc2=nn.Linear(hidden_dim,hidden_dim//2)self.fc3=nn.Linear(hidden_dim//2,output_dim)defforward(self,x):x=F.relu(self.fc1(x))x=F.relu(self.fc2(x))x=self.fc3(x)returnx# Define model parametersinput_dim=X_train.shape[1]hidden_dim=256output_dim=len(np.unique(y_train))model=SimpleCNN(input_dim,hidden_dim,output_dim)

训练代码

接下来,我们将编写训练代码以训练我们的模型。

[<title="Training Code for CNN on Hyperspectral Data">]importtorch.optimasoptimfromsklearn.metricsimportclassification_report,accuracy_score# Define loss function and optimizercriterion=nn.CrossEntropyLoss()optimizer=optim.Adam(model.parameters(),lr=0.001)# Training loopnum_epochs=100device=torch.device("cuda"iftorch.cuda.is_available()else"cpu")model.to(device)forepochinrange(num_epochs):model.train()running_loss=0.0forinputs,labelsintrain_loader:inputs=inputs.float().to(device)labels=labels.long().to(device)optimizer.zero_grad()outputs=model(inputs)loss=criterion(outputs,labels)loss.backward()optimizer.step()running_loss+=loss.item()avg_loss=running_loss/len(train_loader)print(f"Epoch [{epoch+1}/{num_epochs}], Train Loss:{avg_loss:.4f}")# Evaluation on validation setmodel.eval()correct=0total=0all_preds=[]all_labels=[]withtorch.no_grad():forinputs,labelsintest_loader:inputs=inputs.float().to(device)labels=labels.long().to(device)outputs=model(inputs)_,predicted=torch.max(outputs.data,1)total+=labels.size(0)correct+=(predicted==labels).sum().item()all_preds.extend(predicted.cpu().numpy())all_labels.extend(labels.cpu().numpy())accuracy=correct/totalprint(f"Epoch [{epoch+1}/{num_epochs}], Test Accuracy:{accuracy:.4f}")# Print classification reportprint(classification_report(all_labels,all_preds,target_names=[str(i)foriinrange(output_dim)]))

模型评估

在训练过程中,计算了准确率和其他指标。为了更好地可视化模型性能,我们可以绘制混淆矩阵。

[<title="Confusion Matrix Visualization">]fromsklearn.metricsimportconfusion_matriximportseabornassnsimportmatplotlib.pyplotasplt# Compute confusion matrixconf_mat=confusion_matrix(all_labels,all_preds)# Plot confusion matrixplt.figure(figsize=(20,16))sns.heatmap(conf_mat,annot=True,fmt='d',cmap='Blues',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数据集进行训练和评估。以下是所有相关的代码文件:

  1. 数据加载和预处理(load_and_preprocess.py)
  2. CNN模型实现(simple_cnn.py)
  3. 训练代码(training_code.py)
  4. 混淆矩阵可视化(confusion_matrix.py)

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询