基于PyTorch的图像分类特征提取与模型训练文档

概述

本代码实现了一个基于PyTorch的图像特征提取与分类模型训练流程。核心功能包括:

  1. 使用预训练ResNet18模型进行图像特征提取

  2. 将提取的特征保存为标准化格式

  3. 基于提取的特征训练分类模型

代码结构详解

1. 库导入

python 复制代码
import torch
import torch.nn as nn
import torchvision
from torchvision import transforms, datasets
from torch.utils.data import DataLoader, Subset
import numpy as np
import os
from ml.model_trainer import ModelTrainer
  • 关键库说明

    • torch:PyTorch核心库

    • torch.nn:神经网络模块

    • torchvision:计算机视觉专用模块

    • numpy:数值计算库

    • os:文件系统操作

    • ModelTrainer:自定义模型训练类(需另行实现)

2. 特征提取器类(FeatureExtractor)

初始化方法 __init__
python 复制代码
def __init__(self):
    self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    self.model = torchvision.models.resnet18(weights='IMAGENET1K_V1')
    self.model = nn.Sequential(*list(self.model.children())[:-1])
    self.model = self.model.to(self.device).eval()
    self.transform = transforms.Compose([...])
  • 功能说明

    • 设备检测:自动选择GPU/CPU

    • 模型加载:使用ImageNet预训练的ResNet18

    • 模型修改:移除最后的全连接层(保留卷积特征提取器)

    • 预处理设置:标准化图像尺寸和颜色空间

特征提取方法 extract_features
python 复制代码
def extract_features(self, data_dir):
    full_dataset = datasets.ImageFolder(...)
    loader = DataLoader(...)
    
    features = []
    labels = []
    with torch.no_grad():
        for inputs, targets in loader:
            inputs = inputs.to(self.device)
            outputs = self.model(inputs)
            features.append(outputs.squeeze().cpu().numpy())
            labels.append(targets.numpy())
    
    features = np.concatenate(...)
    labels = np.concatenate(...)
    return features, labels, full_dataset.classes
  • 关键参数

    • data_dir:包含分类子目录的图像数据集路径

    • batch_size=32:平衡内存使用与处理效率

    • num_workers=4:多线程数据加载

  • 处理流程

    1. 创建ImageFolder数据集

    2. 使用DataLoader批量加载

    3. 禁用梯度计算加速推理

    4. 特征维度压缩(squeeze)

    5. 设备间数据传输(GPU->CPU)

    6. 合并所有批次数据

3. 主执行流程

参数配置
python 复制代码
DATA_DIR = "/home/.../data"  # 实际数据路径
SAVE_PATH = "./features.npz"  # 特征保存路径
特征提取与保存
python 复制代码
extractor = FeatureExtractor()
if not os.path.exists(SAVE_PATH):
    features, labels, classes = extractor.extract_features(DATA_DIR)
    np.savez(SAVE_PATH, features=features, labels=labels, classes=classes)
else:
    data = np.load(SAVE_PATH)
    features = data['features']
    labels = data['labels']
  • 文件结构

    • features: N_samples, 512 的特征矩阵

    • labels: N_samples 的标签数组

    • classes: 类别名称列表

模型训练与保存
python 复制代码
X, y = features, labels
trainer = ModelTrainer()
model = trainer.train_model(X, y)
joblib.dump(model, 'pest_classifier.pkl')
  • 假设条件

    • ModelTrainer需实现训练逻辑(如SVM、随机森林等)

    • 默认使用全部数据进行训练(建议实际添加数据分割)

技术细节说明

1. 图像预处理流程

2. 特征维度分析

  • ResNet18最后层输出:512维特征向量

  • 假设1000张图像:

    • 原始图像:1000×3×224×224 (约150MB)

    • 提取特征:1000×512 (约2MB) → 显著降维

3. 性能优化策略

  • GPU加速:自动检测CUDA设备

  • 批量处理:32张/批平衡效率与内存

  • 缓存机制:避免重复特征提取

  • 梯度禁用:减少内存消耗

相关推荐
郑州光合科技余经理8 分钟前
本地生活服务系统:成品模块和定制接口怎么划界
java·前端·人工智能·后端·系统架构·php·ai编程
2603_9651481113 分钟前
家居百货蓝海:API挖掘高复购率生活小商品
大数据·服务器·人工智能·python·生活
2501_9304724424 分钟前
踩坑|CodeBuddy权限配置:AI误删文件、乱跑命令、.env泄露怎么防
人工智能·ai编程
小淮AI3 小时前
从“刷题”到“追问”:AI课堂正在重塑哪些学习旧习惯?
人工智能·学习
cui_ruicheng4 小时前
LangChain 应用开发(十四):Agent 上下文与记忆机制
服务器·人工智能·python·langchain
飞哥数智坊4 小时前
直播半小时,我聊了聊 Agent 最常见的 3 个疑问
人工智能
ggb喔5 小时前
AI 原生攻防时代:2026 年渗透测试与逆向开发的范式重构
人工智能·重构
宸津-代码粉碎机5 小时前
AI攻防战升级!基于Spring AI构建Java应用自动免疫安全体系
java·大数据·开发语言·人工智能·python·安全·spring
克里斯蒂亚诺更新5 小时前
机器学习库sklearn的主要任务
人工智能·机器学习·sklearn
先跑起来再说6 小时前
Qavor:一个能跑、能观测、能扩展的开源 AI Agent 工作台
人工智能·开源