资讯动态

多模态垃圾分类系统实战:双塔融合架构与工程落地避坑指南

发布时间:2026/10/9 11:45:22 来源:尧图企业网站定制
简介本资源为基于Python实现的多模态垃圾分类系统完整课程设计资料包面向高校计算机、人工智能相关专业学生及需要完成课设或项目实训的开发者。系统综合利用图像与文本两种模态信息支持可回收物、有害垃圾、厨余垃圾和其他垃圾四类识别涵盖数据采集、预处理、特征提取、分类模型与用户界面等分层架构并附有需求分析与设计文档。压缩包共1282个文件以621个py源码、544个pyc编译文件、32个proto协议定义及若干txt、json、xml配置为主整体约77.96MB目录结构完整便于按模块查阅与二次开发。已有148人学习下载。读者可获得可运行源码、课程设计报告、项目文档及模型相关文件适合作为课设参考、多模态分类入门实践与排错思路借鉴。1. 多模态垃圾分类系统从单张图片到「图文」融合的工程落地很多人做垃圾分类项目第一反应是拿 ResNet 跑一遍 TrashNet 就交差。但真实场景里一张照片能提供的信息非常有限——透明塑料杯和玻璃杯在某个角度下几乎一模一样沾了油渍的外卖盒到底是「可回收」还是「其他垃圾」也常常模棱两可。单模态图像分类的准确率卡在 85% 上下就上不去了这不是模型不够深的问题是信息量本身不够。多模态垃圾分类系统要解决的就是这件事除了图片再引入文本描述比如用户输入「喝完的奶茶杯里面有珍珠」用两个编码器分别提取视觉特征和语义特征再做融合分类。这套方案适合做课程设计、毕业设计也适合想入门多模态融合的开发者——它比图文检索轻量比纯图像分类有技术纵深而且数据集可以自己造。下面从架构选型一路讲到训练、融合、部署和踩坑能直接照着复现。2. 架构选型与数据准备为什么用双塔而不是单塔2.1 双塔融合与单塔拼接的取舍多模态融合常见三条路线早期融合early fusion、晚期融合late fusion、中期融合cross-attention。早期融合就是把图像像素和文本 token 拼在一起送进一个 Transformer听起来优雅但对垃圾分类这种类别少、数据量小的任务来说训练成本高且容易过拟合。晚期融合是各自出 logits 再加权平均实现最简单但两个模态之间没有交互文本分支基本沦为「纠错补丁」。我一般选中期融合里的双塔结构图像走 CNN 或 ViT文本走轻量 Transformer 或 TextCNN各自出 embedding 后做 cross-attention 或简单的门控融合。理由是——两个模态在中间层交互既有信息互补又不会因为参数量爆炸导致小数据集训不动。具体来说图像塔用预训练 ResNet50 或 EfficientNet-B0文本塔用 6 层 Transformerhidden256融合层用 4 头 cross-attention最后接两层 MLP 出分类头。选这个结构的另一个好处是推理时可以只跑图像塔做快速分类文本塔作为可选增强。部署到边缘设备时如果用户没输入文本系统自动降级为单模态不至于整个服务挂掉。2.2 数据集构建图像采集与文本标注公开数据集里TrashNet 只有 6 类共 2527 张图类别少、场景单一。做多模态必须自己补数据。我的做法是分两步走第一步图像数据。用手机在厨房、办公室、小区垃圾桶旁拍每个大类至少 500 张覆盖不同光照、角度、遮挡。拍完用 labelImg 或 Label Studio 标分类标签。注意别只拍「干净样本」——沾油的纸盒、撕掉标签的塑料瓶、混在一起的垃圾堆这些才是真实分布。第二步文本描述。每张图配一句自然语言描述格式不固定但必须包含材质、状态、使用场景三个要素。比如「透明塑料矿泉水瓶已喝完瓶身有标签」或「陶瓷马克杯杯口有缺口放在办公桌上」。文本不需要长15 到 40 字足够。标注时让不同人写保留语言多样性别用模板批量生成——否则文本塔学到的只是模板特征不是真实语义。数据组织成如下目录结构dataset/ ├── images/ │ ├── recyclable/ │ ├── kitchen_waste/ │ ├── hazardous/ │ └── other/ ├── texts/ │ ├── train.json │ ├── val.json │ └── test.json └── splits/ ├── train.txt ├── val.txt └── test.txttrain.json 里每条记录包含 image_path 和 description 两个字段。划分比例按 7:1.5:1.5且必须保证同一场景拍的多张图不跨集——否则验证集准确率虚高上线就翻车。2.3 数据增强与类别不平衡处理垃圾分类天然不平衡可回收物最多有害垃圾最少。直接训练会导致模型偏向多数类。我的处理策略是「增强 重采样 损失加权」三件套。图像增强用 Albumentations随机水平翻转、±15 度旋转、亮度对比度扰动、高斯模糊模拟运动模糊、随机遮挡模拟垃圾被部分遮盖。文本增强用同义词替换和随机删除但删除比例不超过 15%否则语义会变。重采样用 WeightedRandomSampler每个样本的权重设为 1/类别样本数。损失函数用 Focal Lossgamma2让模型关注难分类样本。这三招下来少数类召回率能从 0.4 提到 0.7 以上。import torch from torch.utils.data import WeightedRandomSampler from collections import Counter # 假设 labels 是训练集所有样本的类别列表 class_counts Counter(labels) class_weights {cls: 1.0 / count for cls, count in class_counts.items()} sample_weights [class_weights[label] for label in labels] sampler WeightedRandomSampler( weightssample_weights, num_sampleslen(sample_weights), replacementTrue # 允许重复采样少数类 )这段代码的关键在 replacementTrue它让少数类样本被反复抽到从而在每个 batch 里占比接近均衡。num_samples 设成训练集大小保证一个 epoch 看到的样本数和原来一致。注意别设太大否则少数类过拟合严重。3. 模型实现图像塔、文本塔与融合层的代码落地3.1 图像分支用预训练 EfficientNet 做特征提取图像塔我选 EfficientNet-B0理由是它在 ImageNet 上预训练权重小约 20MB推理快且特征维度1280适中方便和文本特征对齐。如果你追求更高精度可以换 ConvNeXt-Tiny 或 Swin-Tiny但训练显存至少翻倍。import torch.nn as nn from torchvision.models import efficientnet_b0, EfficientNet_B0_Weights class ImageEncoder(nn.Module): def __init__(self, pretrainedTrue, freeze_backboneFalse): super().__init__() weights EfficientNet_B0_Weights.IMAGENET1K_V1 if pretrained else None self.backbone efficientnet_b0(weightsweights) # 去掉原始分类头保留特征提取部分 in_features self.backbone.classifier[1].in_features self.backbone.classifier nn.Identity() self.proj nn.Sequential( nn.Linear(in_features, 512), nn.BatchNorm1d(512), nn.GELU(), nn.Dropout(0.3) ) if freeze_backbone: for param in self.backbone.parameters(): param.requires_grad False def forward(self, x): feat self.backbone(x) # (B, 1280) return self.proj(feat) # (B, 512)这里把 1280 维特征投影到 512 维是为了和文本塔的 hidden 维度对齐方便后续做 cross-attention。BatchNorm1d 在 batch size 小于 16 时统计量不稳定如果显存不够只能用小 batch建议换成 LayerNorm。freeze_backbone 在前 5 个 epoch 设为 True让分类头先 warmup之后再解冻全量微调——这是避免预训练权重被随机初始化的头带偏的常用做法。3.2 文本分支轻量 Transformer 编码中文描述文本塔不用 BERT-base太大。我用 6 层 TransformerEncoderhidden2564 头注意力词表用 jieba 分词后统计训练集构建大小控制在 8000 以内。嵌入层用随机初始化因为垃圾分类的文本描述和通用语料差异大预训练词向量反而可能引入噪声。import torch.nn as nn import math class TextEncoder(nn.Module): def __init__(self, vocab_size, hidden256, nhead4, num_layers6, max_len64): super().__init__() self.embed nn.Embedding(vocab_size, hidden, padding_idx0) self.pos_enc PositionalEncoding(hidden, max_len) encoder_layer nn.TransformerEncoderLayer( d_modelhidden, nheadnhead, dim_feedforwardhidden * 4, dropout0.2, batch_firstTrue ) self.transformer nn.TransformerEncoder(encoder_layer, num_layersnum_layers) self.proj nn.Linear(hidden, 512) def forward(self, input_ids, attention_mask): x self.embed(input_ids) * math.sqrt(self.embed.embedding_dim) x self.pos_enc(x) # 把 padding 位置 mask 掉避免注意力分配到无意义 token x self.transformer(x, src_key_padding_mask(attention_mask 0)) # 取非 padding 位置的平均池化 mask attention_mask.unsqueeze(-1).float() x (x * mask).sum(dim1) / mask.sum(dim1).clamp(min1e-9) return self.proj(x)PositionalEncoding 用标准正弦编码max_len64 足够覆盖 40 字以内的描述。关键在 src_key_padding_mask——如果不传这个参数padding token 会参与注意力计算导致短文本的特征被稀释。平均池化比取 [CLS] 更稳因为我们的 Transformer 没有专门训练 [CLS] token。3.3 融合层Cross-Attention 与门控融合的对比融合层是整个系统最核心的部分。我试过三种方案融合方式参数量验证集准确率训练稳定性拼接 MLP1.2M88.3%高Cross-Attention2.8M91.7%中门控融合1.5M90.5%高Cross-Attention 效果最好但训练时 loss 震荡明显需要 warmup 和梯度裁剪。门控融合是折中方案用一个可学习的门控向量控制两个模态的贡献比例实现简单且稳定。class GatedFusion(nn.Module): def __init__(self, dim512): super().__init__() self.gate nn.Sequential( nn.Linear(dim * 2, dim), nn.Sigmoid() ) self.norm nn.LayerNorm(dim) def forward(self, img_feat, txt_feat): # 门控值决定文本特征保留多少 g self.gate(torch.cat([img_feat, txt_feat], dim-1)) fused img_feat * g txt_feat * (1 - g) return self.norm(fused)门控融合的逻辑是当图像质量差模糊、遮挡时门控值趋近 0系统自动依赖文本特征当文本描述缺失或模糊时门控值趋近 1退化为纯图像分类。这个自适应机制在真实场景里非常实用。3.4 训练循环与关键超参数训练用 AdamWlr3e-4weight_decay0.01cosine 退火warmup 500 步。batch size 设 32图像 224x224文本 64 token。梯度裁剪 max_norm1.0。总共训 50 个 epoch前 5 个 epoch 冻结图像 backbone。optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay0.01) scheduler torch.optim.lr_scheduler.OneCycleLR( optimizer, max_lr3e-4, total_stepsnum_epochs * len(train_loader), pct_start0.1 # 前 10% 步数做 warmup ) criterion FocalLoss(gamma2.0, alphaclass_weights) for epoch in range(num_epochs): model.train() for batch in train_loader: images batch[image].cuda() input_ids batch[input_ids].cuda() mask batch[attention_mask].cuda() labels batch[label].cuda() img_feat image_encoder(images) txt_feat text_encoder(input_ids, mask) logits classifier(fusion(img_feat, txt_feat)) loss criterion(logits, labels) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() scheduler.step() optimizer.zero_grad()OneCycleLR 比 StepLR 收敛快但 pct_start 别设太大0.1 到 0.15 之间比较稳。FocalLoss 的 alpha 传类别权重gamma2 是经验值如果难样本不多可以降到 1。梯度裁剪一定要加Cross-Attention 层在初期容易梯度爆炸。4. 避坑与排查多模态垃圾分类的 5 个血泪教训4.1 文本塔过拟合验证集 loss 先降后升现象训练到第 8 个 epoch文本塔的验证 loss 开始上升但图像塔还在降整体准确率停滞。原因文本描述标注时用了太多重复句式模型记住了模板而不是语义。比如「这是一个XX已经XX」出现频率过高。解决标注时强制要求不同人用不同句式且对高频模板做下采样。另外在文本塔加 Dropout0.3 和权重衰减 0.05比图像塔的 0.01 更激进。4.2 模态失衡图像塔压倒文本塔现象门控值在训练后全部趋近 1文本特征几乎不起作用融合退化成单模态。原因图像塔用了预训练权重初始特征质量远高于随机初始化的文本塔梯度更新时图像塔主导了融合层。解决前 5 个 epoch 冻结图像 backbone只训文本塔和融合层或者给文本塔设更大的学习率图像塔 1e-4文本塔 5e-4。我一般两个都做。4.3 推理时文本缺失导致服务崩溃现象部署后用户不输入文本系统直接报错或输出随机结果。原因训练时每个样本都有文本模型没学过「文本为空」的情况。解决训练时随机将 10% 样本的文本替换为空字符串让模型学会在文本缺失时依赖图像。推理时如果文本为空门控值自动趋近 1走纯图像分支。4.4 图像预处理不一致导致精度骤降现象本地验证准确率 91%部署到服务端后掉到 76%。原因训练用 Albumentations 做归一化mean[0.485,0.456,0.406]部署时用了 OpenCV 默认的 BGR 和 0-255 范围没做一致的归一化。解决把预处理逻辑封装成一个类训练和推理共用同一份代码。别在部署时重写预处理。4.5 类别标签映射错位现象模型预测「有害垃圾」的样本实际是「可回收物」但混淆矩阵显示两类互相错分严重。原因数据标注时类别文件夹按字母排序但标签映射表按中文拼音排序导致索引错位。解决用 JSON 文件显式定义类别到索引的映射训练和推理都读同一个文件。别依赖文件夹遍历顺序。5. 进阶技巧用置信度校准和 TTA 把准确率再提 3 个点模型训完之后别急着交报告。还有两个几乎零成本的技巧能再榨出几个点。置信度校准。多模态融合后模型在某些样本上会过度自信——比如图像模糊但文本描述清晰时softmax 输出 0.99实际错了。用 Temperature Scaling 在验证集上拟合一个温度参数 T把 logits 除以 T 后再 softmax。T 通常取 1.2 到 2.0 之间能让置信度更接近真实准确率。代码就几行class TemperatureScaler(nn.Module): def __init__(self): super().__init__() self.temperature nn.Parameter(torch.ones(1) * 1.5) def forward(self, logits): return logits / self.temperature # 在验证集上优化 temperature 参数 scaler TemperatureScaler().cuda() optimizer torch.optim.LBFGS([scaler.temperature], lr0.01, max_iter50) # ... 用 NLL loss 拟合校准后模型输出的置信度可以直接用来做拒识——低于 0.6 的样本转人工审核避免自动分错。测试时增强TTA。推理时对同一张图做 5 种变换原图、水平翻转、±10 度旋转、中心裁剪分别预测后取平均 logits。文本侧也可以做同义词替换后的多次预测。TTA 能让准确率再提 1.5 到 2.5 个点代价是推理时间翻 5 倍。如果服务对延迟不敏感这个技巧非常划算。我自己的习惯是先把 Temperature Scaling 加上确认置信度可靠后再决定要不要上 TTA。如果业务允许 200ms 以上的延迟TTA 必加如果要求 50ms 以内就只做校准。最后说一个教训别在测试集上调任何参数。我见过有人用测试集选温度 T结果上线后校准完全失效。验证集就是验证集测试集只在最后跑一次。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑