
♂️ 个人主页艾派森的个人主页✍作者简介Python学习者 希望大家多多支持我们一起进步如果文章对你有帮助的话欢迎评论 点赞 收藏 加关注目录1.项目背景2.数据集介绍3.技术工具4.实验过程4.1导入数据4.2数据预处理4.3数据可视化4.4构建模型4.5训练模型4.6模型评估5.总结源代码1.项目背景在自动驾驶、无人机智能巡航、户外安防监控以及防灾减灾等安全关键型自主系统的落地推进中计算机视觉模型在极端恶劣天气条件下的全天候鲁棒性已成为制约技术跨越和工业化上线的核心瓶颈。传统的通用目标检测与识别算法其特征提取能力大多依托于光照良好、视线清晰的常规自然环境一旦切入到暴雨、大雪、浓雾、强霾甚至是龙卷风、野火、洪涝等罕见且兼具巨大破坏力的灾难性极端天气场景时图像往往会因大面积像素退化而出现对比度骤降、边缘轮廓严重模糊以及高频非线性噪声大范围交织的现象导致常规模型出现特征捕获失效和探测精度断崖式下跌的严重偏置。为了从根本上扫除复杂环境引发的领域自适应障碍迫切需要研究如何利用具备更强空间抽象和恒等映射能力的深层残差网络在剧烈干扰中精准剥离出本质的物体几何表型与气象共生特征。本项目立足于这一核心工程痛点依托 2018 至 2026 年间跨度极宽、包含大量长尾与高危灾难场景的真实气象图像资产探索在深度学习框架下构建高性能、多轴向的流式数据预增强防御管线。通过在计算图中深度融合预训练 ResNet50 的强大通用视觉先验并辅以自适应二阶动量优化器和精细化梯度步长对全连接决策头执行定向微调和参数淬炼旨在攻克形态退化场景下的细粒度特征识别难题从而为全天候高可靠性智能视觉监控与边缘端安全应急部署沉淀出极具实战价值的技术底座。2.数据集介绍本实验数据集来源于KaggleXWOD极端天气目标检测是一个开创性的数据集旨在拓展计算机视觉鲁棒性的边界。大多数现有数据集侧重于雨雾等常见天气而XWOD则引入了影响巨大、罕见的灾难性事件包括龙卷风、野火和洪水。该数据集收集于2018年至2026年间为领域自适应和安全关键型自主系统提供了一个严苛的测试平台。背景与灵感XWOD 的灵感源于一个关键的安全差距虽然基于摄像头的感知技术已在乘用车和无人驾驶出租车中得到大规模应用但应对极端天气的能力仍然是人工智能和人类驾驶员之间的“剩余差距”。现有数据凸显了问题的紧迫性-安全危机天气事件约占交通事故的 12% 和交通死亡人数的 9%。-系统故障最近对数百万辆自动驾驶汽车和无人驾驶出租车在雾中“失灵”的调查表明目前的车型还无法应对能见度降低的情况。气候放大效应气候变化正在改变灾害的分布格局。感知系统现在必须应对比十年前更频繁、更强烈的野火烟雾、山洪暴发和强对流事件。XWOD 的创建是为了解决现有数据集中的四个关键限制规模、天气覆盖范围、地理偏差和来源透明度。数据来源数据来自哪里与仅从开放网络抓取的数据集不同XWOD 是对高风险环境的精心综合-历史交通档案记录罕见、影响巨大的事件例如洪水、暴雪和龙卷风事件。-真实行车记录仪画面车辆在龙卷风和暴雨中行驶的实时视角。-公民科学与社交媒体提供来自世界各地不同地点的野火和灾难性天气的原始视频以减轻区域数据集中存在的地理偏差。标准化评估拥抱不平衡我们有意避免“完美”的阶级平衡因为现实世界是不平衡的。-不平衡因素在野火中你会看到汽车比自行车多得多在洪水中卡车可能比行人更容易被看到。-拆分逻辑我们预先拆分的数据集62/15/23保持了这些自然分布。这迫使模型处理长尾先验分布和域偏移。数据集具有最广泛的覆盖范围唯一包含气候加剧灾害如洪水、龙卷风和野火以及雨、雪和雾的数据集。-全球多样性结合北美、亚洲和欧洲交通环境的数据以确保模型的普适性。-专家审核经过人工审核以确保即使在极端遮挡或环境噪声的情况下检测目标仍然有效。视觉概览这些代表性样本展示了 XWOD 基准测试的多样性涵盖了从普通降水到灾难性气候事件的各种情况该数据集提供了各种真实世界的极端天气图像并附有精确的标注。与其他数据集不同XWOD 捕捉到了洪水、龙卷风和野火等罕见且影响巨大的事件从而拓展了计算机视觉技术的边界。3.技术工具Python版本:3.9代码编辑器jupyter notebook4.实验过程4.1导入数据在面对暴雨、大雪、浓雾、强霾等极端天气场景时传统的计算机视觉模型往往由于图像对比度骤降、边缘轮廓模糊以及噪声干扰严重导致探测识别率出现断崖式下跌。在复杂气象条件下实现鲁棒的物体探测首先需要依赖强大的高阶空间特征抽象能力预训练的ResNet50凭借其独特的残差连接结构能够有效抑制深层网络中的梯度消失问题成为应对这类非线性噪声干扰的理想骨干网。在PyTorch框架下展开实验的第一步是精准构建底层环境的依赖闭环并打通物理磁盘上的数据集资产通道。代码首先将硬件算力调度指针动态绑定至GPU设备随后针对多品类的极端天气图像开展字符串前缀扫描与离散映射通过数据结构化算子将其重组为易于管线吞吐的Pandas核心帧从而为后续构建高性能的流式定制化张量分发器沉淀出清晰的索引矩阵。# # 第一部分第三方工业级核心库与算力设备配置 # import os import numpy as np import pandas as pd import matplotlib.pyplot as plt from PIL import Image from tqdm import tqdm import torch import torch.nn as nn from torch.utils.data import Dataset from torch.utils.data import DataLoader from torchvision import transforms from torchvision.models import resnet50 from sklearn.model_selection import train_test_split from sklearn.metrics import classification_report from sklearn.metrics import confusion_matrix # 动态感知硬件环境优先将全局张量计算图挂载至 CUDA 核心以激活硬件加速 device torch.device(cuda if torch.cuda.is_available()else cpu) # # 第二部分物理物理路径配置与天气类目多轴映射 # DATASET_PATH /kaggle/input/datasets/kuantinglai/exwod/dataset TRAIN_IMAGES f{DATASET_PATH}/train/images VALID_IMAGES f{DATASET_PATH}/valid/images TEST_IMAGES f{DATASET_PATH}/test/images # 显式注册极端天气类目的独热离散数值编码标签 weather_map { heavy_rain:0, snow:1, fog:2, haze:3, flooding:4, tornado:5, wildfire:6 } # 异步反转映射表便于后续预测阶段将离散索引快速反解为物理物理文本 idx_to_weather { v:k for k,v in weather_map.items() } # # 第三部分文件前缀自检与全盘资产结构化归集 # files os.listdir(TRAIN_IMAGES) data [] # 线性流扫描物理磁盘目录下的图像文件 for file in files: label None # 动态匹配文件名字符串前缀锁定其所属的极端气象类目 for weather in weather_map: if file.startswith(weather): label weather_map[weather] break # 若标签解析合法则将物理绝对路径与离散标签打包存根 if label is not None: data.append([ os.path.join(TRAIN_IMAGES,file), label ]) # 依托 Pandas 画布构建标准结构化二维数据帧 df pd.DataFrame( data, columns[ image_path, label ] ) # 动态回显顶端前 5 行矩阵样本进行路径对齐自检 df.head()本环节成功实现了物理图像资产向深度学习标准工程数据帧的平稳过渡。代码利用前缀匹配逻辑file.startswith直接在底层文件系统层面完成了对 heavy_rain、snow、fog 等 7 类复杂极端气象特征的标签解耦从源头上建立起了一种高度松耦合的映射关系。通过最后执行的df.head()自检命令控制台同步呈现的结构化数据帧不仅清晰反映了各图像资产物理存储位置与对应数值标签的绑定状态更彻底消除了由于文件乱序对后续矩阵化批处理可能产生的干扰。这一规范的数据清洗形态为下一阶段构建多线程异步流式图像加载管线Dataset/DataLoader提供了确定性的静态索引支撑。4.2数据预处理在极端天气图像的识别任务中光照突变、雨雪造成的视线遮挡以及雾霾引起的对比度退化都会使模型在提取底层特征时产生严重偏置。为了从根本上增强网络对多变气象条件的泛化抗扰性本阶段首先利用分层采样策略对大盘数据实施科学切分确保训练集与验证集在各大恶劣天气类别上的样本权重绝对对齐。随后代码在训练流水线上注入了定制化的空间与色彩色彩扰动算子通过随机水平翻转、小角度旋转以及对亮度与对比度的动态微调在线模拟现实环境中错综复杂的自然光影变幻而对于验证管线则仅执行标准的尺寸归一化以确保存量指标评估的绝对纯净。最终通过重写 PyTorch 的定制化数据集容器并挂载多线程数据加载器实现了物理磁盘资产向显存高并发异步张量流的底层蜕变。# # 第一部分分层切分数据集严防类别分布倾斜 # # 引入 stratify 约束确保训练集与验证集内部 7 种天气变种的样本比例与原始大盘严格一致 train_df, val_df train_test_split(df,test_size0.2,random_state42,stratifydf[label]) # # 第二部分定制化多轴数据增强与增强与张量规范化算子链 # # 训练集增强链通过空间与色彩扰动人工模拟恶劣气象条件下的复杂环境变幻 train_transform transforms.Compose([ transforms.Resize((224,224)), # 强制统一图像物理分辨率契合 ResNet50 的输入拓扑 transforms.RandomHorizontalFlip(), # 随机执行水平镜像翻转扩充空间平移不变性先验 transforms.RandomRotation(10), # 允许 ±10° 的微小旋转扰动模拟车载或监控镜头的颠簸 transforms.ColorJitter( brightness0.2, # 动态微调亮度模拟暴雨或浓雾引发的光线昏暗 contrast0.2 # 动态微调对比度模拟强霾场景下的视线能见度衰减 ), transforms.ToTensor() # 破坏 PIL 物理结构归一化压缩至 [0.0, 1.0] 的三维标准浮点张量 ]) # 验证集变换链剥离所有随机扰动仅执行纯粹的物理尺寸对齐 val_transform transforms.Compose([ transforms.Resize((224,224)), transforms.ToTensor() ]) # # 第三部分重构 PyTorch 数据集白盒打通多线程高并发分发 # class WeatherDataset(Dataset): def __init__(self,dataframe,transformNone): # 显式执行索引重置并丢弃历史存根防止多线程并行寻址时产生越界空指针 self.df dataframe.reset_index(dropTrue) self.transform transform def __len__(self): return len(self.df) def __getitem__(self, idx): # 依托位置索引线性检索物理路径与目标编码标签 img_path self.df.iloc[idx][image_path] label self.df.iloc[idx][label] # 强制以 RGB 三通道彩色模式载入图像抹平由于灰度图或者异常四通道带来的输入维度冲突 image Image.open(img_path).convert(RGB) # 动态触发上述定义的数据流变换链 if self.transform: image self.transform(image) return image, label # 实例化训练与验证数据流载体 train_ds WeatherDataset(train_df,train_transform) val_ds WeatherDataset(val_df,val_transform) # 挂载 DataLoader 算子配置 batch_size32 的分批吞吐流激活 num_workers2 执行多线程异步并行预加载 train_loader DataLoader(train_ds,batch_size32,shuffleTrue,num_workers2) val_loader DataLoader(val_ds,batch_size32,shuffleFalse,num_workers2)本环节构建的数据预处理管线从架构层面锁死了数据在训练过程中的高并发吞吐效率。在WeatherDataset类的底层实现中通过在初始化阶段执行reset_index(dropTrue)彻底消除了由于train_test_split乱序切分导致的索引断层这是保障多线程异步寻址安全的底层基石。而在末端实例化的DataLoader算子中shuffleTrue与num_workers2的协同配合使得 CPU 在后台能够利用多独立核心并行完成图像的解码、色彩扰动与张量转化并将组装好的 32 容量的小批次矩阵常驻在缓冲区中。这种流水线式的设计避免了 GPU 算力在等待 I/O 耗时产生的空转为下一阶段将连续的高阶极端天气特征图源源不断地送入 ResNet50 骨干网络做好了完备的工程保障。4.3数据可视化在将洗净的张量流源源不断地倾注给残差模型之前通过可视化算子对增强后的图像资产实施全盘抽检是确保数据增强边界合理性、规避空间畸变偏置的关键步骤。在复杂环境物体的探测任务中过度的几何翻转或色彩压缩可能会剥离掉气象特征本身的数理物理属性例如过度的对比度拉伸可能会使“薄雾”直接退化为“浓霾”。本阶段通过在 Matplotlib 画布上构建 2 x 3 的六宫格随机抽检矩阵直接在数据流的分发源头捕捉批次张量不仅能直观验证水平镜像、微小旋转以及光影抖动算子的实际覆盖效果更能以最直观的视觉反馈检验数据清洗质量防止训练逻辑陷入盲目的黑盒拟合状态。# # 第一部分构建 2x3 六宫格流式自检画布 # fig, axes plt.subplots(2,3,figsize(12,8)) # # 第二部分多轴随机寻址与通道重排跨维渲染 # for ax in axes.flat: # 在当前训练集总池内随机抽取一个物理位置索引 idx np.random.randint(len(train_ds)) # 异步触发定制化 Dataset 中的 __getitem__ 算子动态提取增强变换后的图像与类别离散标签 image, label train_ds[idx] # 核心张量重排将 PyTorch 标准的 (Channels, Height, Width) 格式转换为 Matplotlib 要求的 (Height, Width, Channels) 格式 image image.permute(1,2,0) # 将跨维处理后的张量矩阵投射至对应的子图画布中 ax.imshow(image) # 通过前期建立的 idx_to_weather 反向解析映射表动态还原并将气象大类英文标识悬挂在子图顶端 ax.set_title(idx_to_weather[label]) # 彻底隐去二维空间像素刻度最大化提高恶劣气象视觉特征的呈现度 ax.axis(off) # # 第三部分刷新多轴视窗输出全盘数据抽检报告 # plt.show()4.4构建模型在打通了异步流式数据分发管线并验证了增强策略的合理性后项目正式切入核心拓扑网络架设阶段。针对极端天气场景下物体边缘退化与多维非线性噪声交织的特质传统的浅层卷积网络往往会因为感受野受限或深层梯度耗尽导致模型失去对宏观表型的解析能力。本阶段直接引入了在超大规模自然视觉数据集上淬炼出深厚通用边缘感知的经典ResNet50残差架构。通过复用其底层的恒等映射Identity Mapping残差块网络能在保持已有高阶语义抓取能力的同时完全免受深层退化问题的困扰。代码利用 PyTorch 的动态计算图机制精准切断其原有的千分类全连接头并将其量身替换为面向七大恶劣天气物体探测任务的定制化全连接决策算子配合具备强正则约束的 AdamW 优化器全盘激活模型的参数进化能效。# # 第一部分调取经典残差骨干并执行末端决策头定向重构 # # 加载预训练 ResNet50 权重资产DEFAULT 算子会自动锁定官方当前最优的特征提取先验 model resnet50(weightsDEFAULT) # 精准拦截原有全连接层fc的输入特征维度并定向改写为 7 类极端天气类目的输出拓扑 model.fc nn.Linear(model.fc.in_features,7) # # 第二部分算力资产流转与核心数理损失算子挂载 # # 将完整的网络拓扑骨架整体推向预先绑定的 GPU/CPU 显存计算节点 model model.to(device) # 挂载标准多分类交叉熵损失算子内部集成了对置信度的 Softmax 归一化和负对数似然计算 criterion nn.CrossEntropyLoss() # # 第三部分权重衰减优化器注入与精细梯度步长配置 # # 部署高阶 AdamW 优化算子注入 1e-4 的稳健初始步长并自动结合 L2 正则化抑制过拟合 optimizer torch.optim.AdamW(model.parameters(),lr1e-4)本环节成功实现了通用计算机视觉先验与特定复杂气象物体探测任务的架构绑定。通过执行model.fc nn.Linear(model.fc.in_features, 7)代码在保留 ResNet50 前 49 层经典卷积块、最大池化和全局平均池化层的前提下仅对最后的密集连接层进行了手术刀式的重构。同时这里选用了现代深度学习中备受推崇的AdamW优化器替代传统的 Adam其核心技术优势在于将解耦后的权重衰减Weight Decay直接应用于梯度更新公式中。由于本项目的目标是让网络在伴随剧烈雨雪噪声的图像中精准剥离物体轮廓1e-4 的精细学习率配合 AdamW 的自适应二阶动量机制能确保骨干网络的数千万个预训练参数在反向传播中实现平滑微调既能快速逼近全局最优点又从数理根源上筑牢了防范参数跑飞或过拟合的工程防线。4.5训练模型在建立了残差骨架网与定制化全连接决策头后项目正式进入了多轮次迭代的计算图参数微调阶段。在极端天气物体探测的实战工程中网络必须通过反复的前向传播与反向传播演进使得全连接层的参数矩阵能够敏锐识别被恶劣气象条件扭曲后的物体轮廓。本阶段通过显式编写标准的 PyTorch 模型训练循环将整个拟合过程置于高强度监管的状态。代码设计了独立的单批次硬性准确率算子在总设 10 个 Epoch 的训练流水线上控制模型在每个周期起始时切入严格的训练模式并通过多线程生成器异步吞吐批处理张量实现梯度的平滑清空、误差的动态反向追溯以及权重的自适应微调为全盘实验的参数进化夯实了底座。# # 第一部分编写自定义单批次命中率计算算子 # def accuracy(outputs, labels): # 沿着通道轴提取置信度最高的最大值索引转化为离散判定标签 preds outputs.argmax(1) # 统计当前批次中预判离散标签与物理真实标签完全重合的样本频数 correct (preds labels).sum().item() # 返回当前小批次在训练流中的即时命中概率 return correct / len(labels) # # 第二部分端到端 10 周期多线程拟合迭代流水线 # epochs 10 for epoch in range(epochs): # 将模型显式设定为训练模式激活可训练状态并开启 Dropout/Batch Normalization 的训练行为 model.train() running_loss 0 running_acc 0 # 挂载 tqdm 进度条算子实时监控流式多线程异步批处理管线的吞吐状态 for images, labels in tqdm(train_loader): # 将当前批次的数据矩阵和类别标签同步流转至绑定的 GPU 算力显存节点 images images.to(device) labels labels.to(device) # 梯度归零清空上一个批次残留在反向传播管线中的参数梯度严防误差累加 optimizer.zero_grad() # 前向传播将物理张量推入残差计算图产出 7 分类的对数几率向量 outputs model(images) # 误差解构计算当前预判分布与物理真实独热映射之间的多分类交叉熵损耗 loss criterion(outputs,labels) # 反向传播驱动误差损耗反向流经 50 层残差图自动演算各层权重的偏导数矩阵 loss.backward() # 参数刷新驱动 AdamW 优化器结合一阶动量与解耦权重衰减对网络参数执行自适应微调 optimizer.step() # 标量累加动态回显并累加当前批次的 Loss 标量以及 Accuracy 百分比 running_loss loss.item() running_acc accuracy(outputs,labels) # 每个 Epoch 周期结束时通过计算全盘 Loader 的算术平均值实时复盘打印当前演进报告 print( fEpoch {epoch1}/{epochs} f Loss{running_loss/len(train_loader):.4f} f Acc{running_acc/len(train_loader):.4f} )4.6模型评估在完成多轮次迭代的参数微调后将模型推向未参与训练的独立验证集进行全面盘点是验证残差模型泛化能力与抗噪强度的标准流程。在复杂气象物体探测任务中由于雨雪雾霾等外界环境干扰极易导致图像底层特征出现高度交织单凭一个宏观的总体命中率很难真正复盘算法的核心瓶颈。本阶段评估工作严格分为流式无偏置置信度提取与多类目细粒度分类报告Classification Report解构两个部分。代码通过切断底层的梯度追踪链并冻结特定的正则统计量以流的形式平滑吞吐全盘验证张量不仅能精准核算出大盘的最终准确率更通过精确量化各独立天气类别下的精确率、召回率与 F1 综合得分全方位透视模型在应对特定恶劣环境时的真实技术身位。# # 第一部分切换无偏置评估模式异步收集离散预测标签 # # 显式冻结 Batch Normalization 的运行时均值与方差统计并锁死 Dropout 的失活行为 model.eval() predictions [] targets [] # 强制切断 PyTorch 底层的动态计算图追踪链全面释放反向传播所需的显存空间加速推理进程 with torch.no_grad(): # 线性遍历多线程异步验证集加载管线 for images, labels in val_loader: # 流转特征矩阵至预绑定的算力硬件节点 images images.to(device) # 前向推理获取当前验证批次在 7 分类空间下的原始对数几率向量 outputs model(images) # 提取概率最高的通道索引作为离散判定预测值 preds outputs.argmax(1) # 利用 extend 算子将当前批次的张量平滑回传至显卡内存并追加至全局存根列表中 predictions.extend(preds.cpu().numpy()) targets.extend(labels.numpy()) # # 第二部分多维硬性指标解构与分类报告渲染 # # 将存根列表转化为标准 NumPy 二维阵列便于矩阵化高效对位比对 predictions np.array(predictions) targets np.array(targets) # 演算大盘总体平均命中率Accuracy acc (predictions targets).mean() print(Validation Accuracy:,acc) # 一键激活工业级分类评估报告全面透视各离散天气的多维泛化表现 print( classification_report( targets, predictions ) )本环节输出的细粒度量化报告以无可辩辩驳的数理数据强力宣告了预训练 ResNet50 在恶劣气象目标探测任务中的压倒性泛化优势。在控制台回显中总体验证集准确率稳稳飙升至99.29%的高位充分证明了前期在多线程数据增强管线中注入的亮度、对比度与小角度旋转等扰动防御不仅没有带偏模型反而成功淬炼出了网络底层应对图像退化的鲁棒性。进一步深度剖析classification_report可以看到在类别 1大雪、类别 4洪涝以及类别 2浓雾等常备细粒度场景下精确率Precision与召回率Recall几乎全线拉满至 1.00表现极其强悍而类别 0暴雨虽略受噪声干扰导致召回率轻微下摆至 0.95但在 1132 张独立测试大盘的高强度检阅下宏观宏观宏观多分类 F1 综合跑分Macro Avg依然锁死在 0.99 的超高水准。这组扎实的数据反馈直接印证了“残差恒等连接迁移微调 fc”策略在复杂工业现场落地的极高实战价值。5.总结本实验基于涵盖 2018 年至 2026 年间真实气象资产的 XWOD极端天气目标检测开创性数据集针对暴雨、大雪、浓雾、强霾、洪涝、龙卷风、野火等影响巨大且罕见的灾难性自然环境物体成功构建并验证了基于预训练 ResNet50 残差架构的深度学习识别系统。在 PyTorch 框架的多线程异步流式吞吐管线支撑下模型通过多轴空间与色彩扰动算子的合理干预成功克服了恶劣场景下图像对比度暴跌、轮廓边缘高度模糊的技术难题。在历经 10 个 Epoch 的高强度权重微调与 AdamW 优化器的自适应拟合演进后网络在完全隔离的 1132 张独立验证集样本上跑出了高达 99.29% 的大盘总体命中率。纵观细粒度分类报告的底层数据流模型在大雪类别1、洪涝类别4等关键天气场景下展现出了完美命中1.00得分的绝对技术身位而在野火类别6和龙卷风类别5等复杂长尾类目中也维持了 0.99 的超高 F1 综合跑分全盘宏观平均Macro Avg锁死在 0.99。这一扎实的拟合形态和极高的特异性表现不仅有力地印证了预训练残差连接网络在应对复杂非线性噪声干扰时的通用先验红利更为自动驾驶、智能安防监控等安全关键型自主系统在非平稳全天候环境下的高难度领域自适应沉淀出了一套兼顾极速推断吞吐与极致鲁棒性的工业级交付标杆。源代码import os import numpy as np import pandas as pd import matplotlib.pyplot as plt from PIL import Image from tqdm import tqdm import torch import torch.nn as nn from torch.utils.data import Dataset from torch.utils.data import DataLoader from torchvision import transforms from torchvision.models import resnet50 from sklearn.model_selection import train_test_split from sklearn.metrics import classification_report from sklearn.metrics import confusion_matrix device torch.device(cuda if torch.cuda.is_available()else cpu) DATASET_PATH /kaggle/input/datasets/kuantinglai/exwod/dataset TRAIN_IMAGES f{DATASET_PATH}/train/images VALID_IMAGES f{DATASET_PATH}/valid/images TEST_IMAGES f{DATASET_PATH}/test/images weather_map { heavy_rain:0, snow:1, fog:2, haze:3, flooding:4, tornado:5, wildfire:6 } idx_to_weather { v:k for k,v in weather_map.items() } files os.listdir(TRAIN_IMAGES) data [] for file in files: label None for weather in weather_map: if file.startswith(weather): label weather_map[weather] break if label is not None: data.append([ os.path.join(TRAIN_IMAGES,file), label ]) df pd.DataFrame( data, columns[ image_path, label ] ) df.head() train_df, val_df train_test_split(df,test_size0.2,random_state42,stratifydf[label]) train_transform transforms.Compose([ transforms.Resize((224,224)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(10), transforms.ColorJitter( brightness0.2, contrast0.2 ), transforms.ToTensor() ]) val_transform transforms.Compose([ transforms.Resize((224,224)), transforms.ToTensor() ]) class WeatherDataset(Dataset): def __init__(self,dataframe,transformNone): self.df dataframe.reset_index(dropTrue) self.transform transform def __len__(self): return len(self.df) def __getitem__(self, idx): img_path self.df.iloc[idx][image_path] label self.df.iloc[idx][label] image Image.open(img_path).convert(RGB) if self.transform: image self.transform(image) return image, label train_ds WeatherDataset(train_df,train_transform) val_ds WeatherDataset(val_df,val_transform) train_loader DataLoader(train_ds,batch_size32,shuffleTrue,num_workers2) val_loader DataLoader(val_ds,batch_size32,shuffleFalse,num_workers2) fig, axes plt.subplots(2,3,figsize(12,8)) for ax in axes.flat: idx np.random.randint(len(train_ds)) image, label train_ds[idx] image image.permute(1,2,0) ax.imshow(image) ax.set_title(idx_to_weather[label]) ax.axis(off) plt.show() model resnet50(weightsDEFAULT) model.fc nn.Linear(model.fc.in_features,7) model model.to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(),lr1e-4) def accuracy(outputs, labels): preds outputs.argmax(1) correct (preds labels).sum().item() return correct / len(labels) epochs 10 for epoch in range(epochs): model.train() running_loss 0 running_acc 0 for images, labels in tqdm(train_loader): images images.to(device) labels labels.to(device) optimizer.zero_grad() outputs model(images) loss criterion(outputs,labels) loss.backward() optimizer.step() running_loss loss.item() running_acc accuracy(outputs,labels) print( fEpoch {epoch1}/{epochs} f Loss{running_loss/len(train_loader):.4f} f Acc{running_acc/len(train_loader):.4f} ) model.eval() predictions [] targets [] with torch.no_grad(): for images, labels in val_loader: images images.to(device) outputs model(images) preds outputs.argmax(1) predictions.extend(preds.cpu().numpy()) targets.extend(labels.numpy()) predictions np.array(predictions) targets np.array(targets) acc (predictions targets).mean() print(Validation Accuracy:,acc) print( classification_report( targets, predictions ) )资料获取更多粉丝福利关注下方公众号获取