遥感图像分类是计算机视觉与地理信息科学的重要交叉领域,随着深度学习技术的快速发展,基于深度学习的遥感图像分类方法正在彻底改变传统的人工解译模式。本文将手把手带你从零搭建一个完整的遥感图像分类项目,涵盖环境配置、数据预处理、模型构建、训练优化到结果分析的全流程,无论你是刚接触深度学习的新手,还是需要完成毕设的学生,都能通过本文掌握实战技能。
1. 遥感图像分类的核心概念
1.1 什么是遥感图像分类
遥感图像分类是指利用计算机算法对遥感影像中的地物进行自动识别和归类的过程。与传统图像分类相比,遥感图像具有多光谱、高空间分辨率、大尺度覆盖等特点,能够识别农田、建筑、水体、森林等不同地物类型。
在实际应用中,遥感图像分类可以用于国土资源调查、环境监测、灾害评估、城市规划等多个领域。随着高分辨率遥感卫星的普及,每天产生的海量遥感数据迫切需要高效的自动分类方法。
1.2 深度学习在遥感图像分类中的优势
传统的遥感图像分类方法主要依赖人工设计的特征提取器,如纹理特征、形状特征等,但这些方法在复杂场景下的泛化能力有限。深度学习通过卷积神经网络(CNN)自动学习图像的特征表示,具有以下显著优势:
- 特征学习自动化:CNN能够从原始像素中自动学习层次化的特征表示,无需人工设计特征提取器
- 高精度分类:在大规模数据集上训练的深度学习模型能够达到接近甚至超过人类专家的分类精度
- 多尺度信息融合:通过不同层级的卷积操作,模型可以同时捕获局部细节和全局上下文信息
- 端到端学习:从原始输入到最终分类结果,整个流程可以统一优化,减少误差累积
1.3 常用深度学习模型对比
在遥感图像分类任务中,常用的深度学习模型包括CNN、RNN和Transformer等。CNN由于其出色的空间特征提取能力,成为遥感图像分类的主流选择:
- CNN:擅长处理网格状数据,通过卷积核滑动提取局部特征,适合图像空间信息建模
- RNN:主要用于序列数据,在遥感时序分析中有一定应用,但计算复杂度较高
- Transformer:近年来在计算机视觉领域表现突出,特别适合建模长距离依赖关系
对于大多数遥感图像分类任务,建议从CNN模型开始,如ResNet、VGG等经典架构,这些模型在准确性和计算效率之间取得了良好平衡。
2. 环境准备与工具配置
2.1 硬件与软件要求
进行深度学习遥感图像分类需要适当的计算资源,以下是推荐配置:
硬件要求:
- GPU:NVIDIA GTX 1060 6GB或更高(建议RTX 3060及以上)
- 内存:16GB RAM(处理大尺寸图像时建议32GB)
- 存储:至少50GB可用空间(用于存储数据集和模型)
软件环境:
- 操作系统:Ubuntu 18.04+ / Windows 10 / macOS
- Python 3.8+(本文使用Python 3.9)
- CUDA 11.3+(GPU加速训练必需)
- cuDNN 8.2+(深度学习库优化)
2.2 深度学习框架安装
PyTorch是目前最流行的深度学习框架之一,具有良好的灵活性和易用性。以下是完整的安装步骤:
# 创建虚拟环境(推荐) conda create -n remote_sensing python=3.9 conda activate remote_sensing # 安装PyTorch及相关依赖 pip install torch==1.12.1+cu113 torchvision==0.13.1+cu113 torchaudio==0.12.1 --extra-index-url https://download.pytorch.org/whl/cu113 # 安装图像处理库 pip install opencv-python pillow scikit-image # 安装科学计算库 pip install numpy pandas matplotlib seaborn # 安装遥感数据处理专用库 pip install rasterio gdal earthpy2.3 数据集准备与介绍
本文使用UC Merced Land Use Dataset,这是一个常用的遥感图像分类基准数据集,包含21个土地覆盖类别,每类有100张256×256像素的图像。
import os import requests import zipfile # 数据集下载函数 def download_ucmerced_dataset(download_path="./data"): os.makedirs(download_path, exist_ok=True) url = "http://weegee.vision.ucmerced.edu/datasets/landuse/images.zip" zip_path = os.path.join(download_path, "uc_merced.zip") # 下载数据集(实际使用时请确保网络连接) print("正在下载UC Merced数据集...") # response = requests.get(url, stream=True) # with open(zip_path, 'wb') as f: # for chunk in response.iter_content(chunk_size=8192): # f.write(chunk) # 解压数据集 # with zipfile.ZipFile(zip_path, 'r') as zip_ref: # zip_ref.extractall(download_path) print("数据集准备完成!") # 调用下载函数 download_ucmerced_dataset()3. 深度学习模型原理与选择
3.1 卷积神经网络基础架构
卷积神经网络是遥感图像分类的核心技术,其基本组成包括:
卷积层:通过滑动窗口提取局部特征,每个卷积核学习不同的特征模式
import torch import torch.nn as nn # 简单的卷积层示例 class SimpleCNN(nn.Module): def __init__(self, num_classes=21): super(SimpleCNN, self).__init__() self.conv1 = nn.Conv2d(3, 32, kernel_size=3, padding=1) self.conv2 = nn.Conv2d(32, 64, kernel_size=3, padding=1) self.pool = nn.MaxPool2d(2, 2) self.fc1 = nn.Linear(64 * 64 * 64, 512) # 根据输入尺寸调整 self.fc2 = nn.Linear(512, num_classes) self.relu = nn.ReLU() self.dropout = nn.Dropout(0.5) def forward(self, x): x = self.pool(self.relu(self.conv1(x))) x = self.pool(self.relu(self.conv2(x))) x = x.view(x.size(0), -1) # 展平 x = self.dropout(self.relu(self.fc1(x))) x = self.fc2(x) return x池化层:降低特征图尺寸,增加平移不变性,减少计算量全连接层:将学习到的特征映射到最终的分类结果
3.2 迁移学习在遥感分类中的应用
对于数据量有限的遥感任务,迁移学习是提升性能的有效策略。我们可以使用在ImageNet上预训练的模型作为特征提取器:
import torchvision.models as models def create_resnet_model(num_classes=21, pretrained=True): """ 创建基于ResNet的迁移学习模型 """ # 加载预训练模型 model = models.resnet50(pretrained=pretrained) # 冻结底层参数(可选) for param in model.parameters(): param.requires_grad = False # 替换最后的全连接层 num_features = model.fc.in_features model.fc = nn.Sequential( nn.Dropout(0.5), nn.Linear(num_features, 512), nn.ReLU(), nn.Dropout(0.3), nn.Linear(512, num_classes) ) return model # 创建模型实例 model = create_resnet_model(num_classes=21) print(f"模型参数量:{sum(p.numel() for p in model.parameters())}")3.3 模型选择策略
根据任务需求选择合适的模型架构:
- 小数据集(<1万张图像):使用轻量级CNN或迁移学习
- 中等数据集(1-10万张图像):ResNet、DenseNet等中等复杂度模型
- 大数据集(>10万张图像):EfficientNet、Vision Transformer等先进模型
对于大多数遥感分类任务,ResNet50在准确性和效率之间提供了良好的平衡。
4. 数据预处理与增强策略
4.1 遥感图像特性分析
遥感图像与自然图像相比具有独特特性,需要在预处理时特别注意:
- 多光谱信息:遥感图像通常包含多个波段(RGB、近红外等)
- 空间分辨率:像素对应实际地理尺寸,影响地物识别精度
- 辐射定标:需要将DN值转换为地表反射率等物理量
- 几何校正:消除传感器姿态、地形等因素引起的形变
4.2 数据预处理流程
完整的数据预处理流程包括读取、标准化和增强:
import torchvision.transforms as transforms from torch.utils.data import Dataset, DataLoader from PIL import Image import os class RemoteSensingDataset(Dataset): def __init__(self, data_dir, transform=None, split='train'): self.data_dir = data_dir self.transform = transform self.split = split self.classes = sorted([d for d in os.listdir(data_dir) if os.path.isdir(os.path.join(data_dir, d))]) self.class_to_idx = {cls_name: i for i, cls_name in enumerate(self.classes)} # 收集图像路径和标签 self.images = [] for class_name in self.classes: class_dir = os.path.join(data_dir, class_name) for img_name in os.listdir(class_dir): if img_name.lower().endswith(('.jpg', '.jpeg', '.png', '.tif')): self.images.append((os.path.join(class_dir, img_name), self.class_to_idx[class_name])) def __len__(self): return len(self.images) def __getitem__(self, idx): img_path, label = self.images[idx] image = Image.open(img_path).convert('RGB') if self.transform: image = self.transform(image) return image, label # 定义数据增强策略 train_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomHorizontalFlip(p=0.5), transforms.RandomVerticalFlip(p=0.5), transforms.RandomRotation(degrees=15), transforms.ColorJitter(brightness=0.2, contrast=0.2), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) val_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ])4.3 针对遥感数据的特殊增强
遥感图像需要特殊的数据增强技术来模拟真实场景变化:
class RemoteSensingAugmentation: """遥感图像专用数据增强类""" @staticmethod def random_crop_with_scale(image, scale_range=(0.8, 1.2)): """随机缩放裁剪,模拟不同分辨率""" pass @staticmethod def spectral_augmentation(image): """光谱增强,模拟不同光照条件""" pass @staticmethod def simulate_atmospheric_effects(image): """模拟大气影响""" pass5. 完整项目实战:土地覆盖分类
5.1 项目架构设计
我们构建一个完整的遥感图像分类系统,包含以下模块:
remote_sensing_classification/ ├── data/ │ ├── raw/ # 原始数据 │ ├── processed/ # 处理后的数据 │ └── splits/ # 训练/验证/测试划分 ├── models/ │ ├── base_model.py # 基础模型定义 │ └── custom_models.py # 自定义模型 ├── utils/ │ ├── data_loader.py # 数据加载工具 │ ├── metrics.py # 评估指标 │ └── visualization.py # 可视化工具 ├── config.py # 配置文件 ├── train.py # 训练脚本 ├── evaluate.py # 评估脚本 └── inference.py # 推理脚本5.2 模型训练实现
下面是完整的模型训练代码:
import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader import time import numpy as np from sklearn.metrics import accuracy_score, confusion_matrix import matplotlib.pyplot as plt class RemoteSensingTrainer: def __init__(self, model, train_loader, val_loader, config): self.model = model self.train_loader = train_loader self.val_loader = val_loader self.config = config self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') self.model.to(self.device) # 优化器和损失函数 self.criterion = nn.CrossEntropyLoss() self.optimizer = optim.Adam(model.parameters(), lr=config['learning_rate']) self.scheduler = optim.lr_scheduler.StepLR(self.optimizer, step_size=config['step_size'], gamma=config['gamma']) # 训练记录 self.train_losses = [] self.val_accuracies = [] self.best_accuracy = 0.0 def train_epoch(self, epoch): self.model.train() running_loss = 0.0 for batch_idx, (data, target) in enumerate(self.train_loader): data, target = data.to(self.device), target.to(self.device) self.optimizer.zero_grad() output = self.model(data) loss = self.criterion(output, target) loss.backward() self.optimizer.step() running_loss += loss.item() if batch_idx % 100 == 0: print(f'Epoch: {epoch} [{batch_idx * len(data)}/{len(self.train_loader.dataset)} ' f'({100. * batch_idx / len(self.train_loader):.0f}%)]\tLoss: {loss.item():.6f}') avg_loss = running_loss / len(self.train_loader) self.train_losses.append(avg_loss) return avg_loss def validate(self, epoch): self.model.eval() val_loss = 0 correct = 0 all_preds = [] all_targets = [] with torch.no_grad(): for data, target in self.val_loader: data, target = data.to(self.device), target.to(self.device) output = self.model(data) val_loss += self.criterion(output, target).item() pred = output.argmax(dim=1, keepdim=True) correct += pred.eq(target.view_as(pred)).sum().item() all_preds.extend(pred.cpu().numpy()) all_targets.extend(target.cpu().numpy()) val_loss /= len(self.val_loader) accuracy = 100. * correct / len(self.val_loader.dataset) self.val_accuracies.append(accuracy) print(f'\nValidation set: Average loss: {val_loss:.4f}, ' f'Accuracy: {correct}/{len(self.val_loader.dataset)} ({accuracy:.2f}%)\n') # 保存最佳模型 if accuracy > self.best_accuracy: self.best_accuracy = accuracy torch.save({ 'epoch': epoch, 'model_state_dict': self.model.state_dict(), 'optimizer_state_dict': self.optimizer.state_dict(), 'accuracy': accuracy }, 'best_model.pth') return accuracy, all_preds, all_targets def train(self): print("开始训练...") for epoch in range(1, self.config['epochs'] + 1): start_time = time.time() train_loss = self.train_epoch(epoch) val_accuracy, _, _ = self.validate(epoch) self.scheduler.step() epoch_time = time.time() - start_time print(f'Epoch {epoch} 完成, 耗时: {epoch_time:.2f}秒') print(f'训练损失: {train_loss:.4f}, 验证准确率: {val_accuracy:.2f}%') # 早停机制 if epoch > 10 and val_accuracy < max(self.val_accuracies[-5:]): print("验证准确率不再提升,提前停止训练") break # 配置参数 config = { 'batch_size': 32, 'learning_rate': 0.001, 'epochs': 50, 'step_size': 10, 'gamma': 0.1 } # 创建数据加载器 train_dataset = RemoteSensingDataset('./data/train', transform=train_transform) val_dataset = RemoteSensingDataset('./data/val', transform=val_transform) train_loader = DataLoader(train_dataset, batch_size=config['batch_size'], shuffle=True) val_loader = DataLoader(val_dataset, batch_size=config['batch_size'], shuffle=False) # 初始化训练器 model = create_resnet_model(num_classes=21) trainer = RemoteSensingTrainer(model, train_loader, val_loader, config) trainer.train()5.3 模型评估与结果分析
训练完成后需要对模型进行全面评估:
def evaluate_model(model, test_loader, class_names): """全面评估模型性能""" model.eval() all_preds = [] all_targets = [] all_probabilities = [] with torch.no_grad(): for data, target in test_loader: data, target = data.to(device), target.to(device) output = model(data) probabilities = torch.softmax(output, dim=1) pred = output.argmax(dim=1) all_preds.extend(pred.cpu().numpy()) all_targets.extend(target.cpu().numpy()) all_probabilities.extend(probabilities.cpu().numpy()) # 计算各项指标 accuracy = accuracy_score(all_targets, all_preds) cm = confusion_matrix(all_targets, all_preds) # 分类报告 from sklearn.metrics import classification_report report = classification_report(all_targets, all_preds, target_names=class_names) print(f"整体准确率: {accuracy:.4f}") print("\n分类报告:") print(report) return accuracy, cm, all_probabilities # 可视化混淆矩阵 def plot_confusion_matrix(cm, class_names): plt.figure(figsize=(12, 10)) plt.imshow(cm, interpolation='nearest', cmap=plt.cm.Blues) plt.title('混淆矩阵') plt.colorbar() tick_marks = np.arange(len(class_names)) plt.xticks(tick_marks, class_names, rotation=45) plt.yticks(tick_marks, class_names) # 添加数值标注 thresh = cm.max() / 2. for i, j in itertools.product(range(cm.shape[0]), range(cm.shape[1])): plt.text(j, i, format(cm[i, j], 'd'), horizontalalignment="center", color="white" if cm[i, j] > thresh else "black") plt.tight_layout() plt.ylabel('真实标签') plt.xlabel('预测标签') plt.show()6. 高级技巧与优化策略
6.1 类别不平衡处理
遥感数据中经常出现类别不平衡问题,需要特殊处理:
# 计算类别权重 from sklearn.utils.class_weight import compute_class_weight def calculate_class_weights(dataset): """计算类别权重用于损失函数""" labels = [label for _, label in dataset.images] class_weights = compute_class_weight('balanced', classes=np.unique(labels), y=labels) return torch.tensor(class_weights, dtype=torch.float32) # 使用加权损失函数 class_weights = calculate_class_weights(train_dataset) criterion = nn.CrossEntropyLoss(weight=class_weights.to(device))6.2 模型集成策略
通过模型集成可以进一步提升分类性能:
class ModelEnsemble: def __init__(self, model_paths, device): self.models = [] for path in model_paths: model = create_resnet_model(num_classes=21) checkpoint = torch.load(path) model.load_state_dict(checkpoint['model_state_dict']) model.to(device) model.eval() self.models.append(model) def predict(self, x): predictions = [] for model in self.models: with torch.no_grad(): output = model(x) prob = torch.softmax(output, dim=1) predictions.append(prob.cpu().numpy()) # 平均概率 avg_prob = np.mean(predictions, axis=0) return np.argmax(avg_prob, axis=1)6.3 超参数优化
使用Optuna等工具进行自动化超参数搜索:
import optuna def objective(trial): # 超参数搜索空间 lr = trial.suggest_float('lr', 1e-5, 1e-2, log=True) batch_size = trial.suggest_categorical('batch_size', [16, 32, 64]) dropout_rate = trial.suggest_float('dropout_rate', 0.1, 0.5) # 创建模型并训练 model = create_resnet_model(num_classes=21) config = {'learning_rate': lr, 'batch_size': batch_size} trainer = RemoteSensingTrainer(model, train_loader, val_loader, config) trainer.train() return trainer.best_accuracy # 执行超参数优化 study = optuna.create_study(direction='maximize') study.optimize(objective, n_trials=50) print(f'最佳超参数: {study.best_params}') print(f'最佳准确率: {study.best_value:.2f}%')7. 实际应用与部署
7.1 单张图像推理
训练好的模型可以用于单张遥感图像的分类:
def predict_single_image(image_path, model, transform, class_names): """对单张图像进行预测""" # 加载图像 image = Image.open(image_path).convert('RGB') original_image = image.copy() # 预处理 input_tensor = transform(image).unsqueeze(0) # 添加batch维度 # 预测 model.eval() with torch.no_grad(): output = model(input_tensor) probabilities = torch.softmax(output, dim=1) predicted_class = output.argmax(dim=1).item() confidence = probabilities[0][predicted_class].item() # 可视化结果 plt.figure(figsize=(10, 5)) plt.subplot(1, 2, 1) plt.imshow(original_image) plt.title('输入图像') plt.axis('off') plt.subplot(1, 2, 2) # 显示类别概率分布 classes = class_names probs = probabilities[0].cpu().numpy() plt.barh(classes, probs) plt.xlabel('概率') plt.title('分类结果') plt.tight_layout() plt.show() print(f'预测类别: {class_names[predicted_class]}') print(f'置信度: {confidence:.4f}') return predicted_class, confidence # 使用示例 class_names = ['agricultural', 'airplane', 'baseballdiamond', 'beach', 'buildings', 'chaparral', 'denseresidential', 'forest', 'freeway', 'golfcourse', 'harbor', 'intersection', 'mediumresidential', 'mobilehomepark', 'overpass', 'parkinglot', 'river', 'runway', 'sparseresidential', 'storagetanks', 'tenniscourt'] # 加载训练好的模型 checkpoint = torch.load('best_model.pth') model.load_state_dict(checkpoint['model_state_dict']) # 对单张图像进行预测 image_path = 'test_image.jpg' predicted_class, confidence = predict_single_image(image_path, model, val_transform, class_names)7.2 批量处理与API部署
对于实际应用,通常需要处理大量图像或提供在线服务:
from flask import Flask, request, jsonify from PIL import Image import io app = Flask(__name__) # 加载模型 model = create_resnet_model(num_classes=21) checkpoint = torch.load('best_model.pth') model.load_state_dict(checkpoint['model_state_dict']) model.eval() @app.route('/predict', methods=['POST']) def predict(): """遥感图像分类API接口""" if 'image' not in request.files: return jsonify({'error': '没有提供图像文件'}), 400 # 读取图像 image_file = request.files['image'] image = Image.open(io.BytesIO(image_file.read())).convert('RGB') # 预处理 input_tensor = val_transform(image).unsqueeze(0) # 预测 with torch.no_grad(): output = model(input_tensor) probabilities = torch.softmax(output, dim=1) predicted_class_idx = output.argmax(dim=1).item() confidence = probabilities[0][predicted_class_idx].item() # 返回结果 result = { 'predicted_class': class_names[predicted_class_idx], 'confidence': confidence, 'all_probabilities': {class_names[i]: float(prob) for i, prob in enumerate(probabilities[0].cpu().numpy())} } return jsonify(result) if __name__ == '__main__': app.run(host='0.0.0.0', port=5000, debug=True)8. 常见问题与解决方案
8.1 训练过程中的典型问题
问题1:过拟合现象严重
- 现象:训练准确率很高,但验证准确率停滞不前
- 解决方案:
- 增加数据增强强度
- 添加更多的Dropout层
- 使用早停机制
- 尝试模型正则化(L1/L2)
# 改进的模型结构,增强正则化 class RegularizedResNet(nn.Module): def __init__(self, num_classes=21, dropout_rate=0.5): super().__init__() self.backbone = models.resnet50(pretrained=True) num_features = self.backbone.fc.in_features # 更强的正则化 self.backbone.fc = nn.Sequential( nn.Dropout(dropout_rate), nn.Linear(num_features, 512), nn.BatchNorm1d(512), nn.ReLU(), nn.Dropout(dropout_rate), nn.Linear(512, num_classes) )问题2:训练损失不下降
- 现象:多个epoch后损失值几乎没有变化
- 解决方案:
- 检查学习率是否合适
- 验证数据预处理是否正确
- 检查模型架构是否合理
- 确认损失函数选择是否正确
8.2 数据相关问题
问题3:类别不平衡导致模型偏向多数类
- 现象:模型对多数类预测准确,但对少数类识别率低
- 解决方案:
- 使用加权损失函数
- 采用过采样或欠采样技术
- 使用Focal Loss等改进的损失函数
# Focal Loss实现 class FocalLoss(nn.Module): def __init__(self, alpha=1, gamma=2, reduction='mean'): super(FocalLoss, self).__init__() self.alpha = alpha self.gamma = gamma self.reduction = reduction def forward(self, inputs, targets): BCE_loss = nn.CrossEntropyLoss(reduction='none')(inputs, targets) pt = torch.exp(-BCE_loss) F_loss = self.alpha * (1-pt)**self.gamma * BCE_loss if self.reduction == 'mean': return torch.mean(F_loss) elif self.reduction == 'sum': return torch.sum(F_loss) else: return F_loss8.3 性能优化问题
问题4:推理速度过慢
- 现象:模型预测单张图像耗时过长
- 解决方案:
- 使用模型量化技术
- 尝试更轻量的模型架构
- 启用GPU加速
- 使用ONNX等优化格式
# 模型量化示例 def quantize_model(model): model.eval() # 动态量化 quantized_model = torch.quantization.quantize_dynamic( model, {nn.Linear, nn.Conv2d}, dtype=torch.qint8 ) return quantized_model # 使用量化模型进行推理 quantized_model = quantize_model(model)9. 最佳实践与工程建议
9.1 数据管理规范
良好的数据管理是项目成功的基础:
- 数据版本控制:使用DVC等工具管理数据集版本
- 数据质量检查:建立自动化的数据质量验证流程
- 元数据管理:完整记录数据的来源、采集时间、预处理方法等信息
- 数据安全:敏感遥感数据需要加密存储和传输
9.2 模型开发流程
建立规范的模型开发流程:
- 探索性数据分析:深入了解数据特性和分布
- 基线模型建立:使用简单模型建立性能基线
- 迭代优化:基于基线逐步改进模型架构和参数
- 交叉验证:使用k折交叉验证确保模型稳定性
- 消融实验:分析各组件对最终性能的贡献
9.3 生产环境部署考虑
将模型部署到生产环境时需要特别注意:
- 模型监控:实时监控模型性能衰减和数据分布变化
- A/B测试:新模型上线前进行充分的A/B测试
- 回滚机制:建立快速回滚到旧版本的机制
- 资源管理:合理分配计算资源,避免资源浪费
9.4 持续学习与模型更新
遥感数据具有时效性,需要建立持续学习机制:
class ContinuousLearning: def __init__(self, model, memory_size=1000): self.model = model self.memory_buffer = [] # 存储历史样本 self.memory_size = memory_size def update_model(self, new_data, new_labels, learning_rate=0.0001): """使用新数据更新模型""" # 将新数据添加到记忆缓冲区 self._update_memory(new_data, new_labels) # 从记忆缓冲区采样进行训练 rehearsal_data, rehearsal_labels = self._sample_from_memory() # 组合新旧数据训练 combined_data = torch.cat([new_data, rehearsal_data]) combined_labels = torch.cat([new_labels, rehearsal_labels]) # 微调模型 self._fine_tune(combined_data, combined_labels, learning_rate) def _update_memory(self, new_data, new_labels): """更新记忆缓冲区""" # 实现记忆管理逻辑 pass def _sample_from_memory(self): """从记忆缓冲区采样""" # 实现采样逻辑 pass通过本文的完整学习,你应该已经掌握了基于深度学习的遥感图像分类从理论到实践的全套技能。在实际项目中,建议先从简单的数据集和模型开始,逐步深入复杂的应用场景。记得始终保持对数据质量的关注,这是影响模型性能的最关键因素。