尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

基于BERT与ResNet50的多模态情感分析:五种融合策略全解析

基于BERT与ResNet50的多模态情感分析:五种融合策略全解析 简介面向多模态情感分析初学者的Python实现项目以BERTResNet50为骨干网络融合文本与图像特征共提供五种融合方案包括两种Naive拼接与三种Attention机制跨模态注意力、输出Transformer编码器等并附有对应模型结构示意图可满足毕设、课程设计或工程实训等场景。资源压缩包共40个文件涵盖17个Python源码、预编译的pyc、依赖清单、JSON数据配置、训练说明文档及融合模型结构图整体仅470KB便于快速部署与二次开发。Models目录按五种融合方法清晰组织配合README和API文件能帮助读者理解不同融合策略的差异与适用场景从数据预处理、模型训练到结果评估均有完整实现适合作为入门多模态研究的参考基线。目前已有558人学习下载适合希望通过对比实验深入掌握多模态情感分析的小白或进阶学习者。1. 先看这套系统在解决什么问题如果只是把文本丢给 BERT、把图片丢给 ResNet各自跑出一个情感极性那其实还停留在“单模态分类”的层面算不上真正的多模态融合。这套项目把文本和图像放在同一个训练流程里用五种不同的融合策略去对比“什么时候该拼接特征”“什么时候该做注意力交互”并且给出了完整的训练、验证、预测代码。它在 MVSA 这类图文配对数据集上做二分类正面/负面也支持换自己的数据。适合三类人一是要做毕设或课程设计、需要一个能跑通且能讲清楚原理的基线系统二是刚接触 BERT 和 ResNet、想直观看到两种预训练模型怎么协同的初学者三是想快速对比融合策略、准备写多模态论文的进阶学习者。项目文件里 Models 目录下五种模型各有独立实现改起来不互相干扰。2. 数据预处理与 BERT ResNet50 的输入构造2.1 数据格式与文本侧清洗项目 data 目录下提供了 train.json、test.json 和 train.txt、test_without_label.txt其中 json 是带标签的结构化数据。可以先看一下 train.json 的字段组织方式import json with open(data/train.json, r, encodingutf-8) as f: samples json.load(f) print(type(samples)) print(samples[0])推荐在命令行或 notebook 里执行python -c import json; djson.load(open(data/train.json)); print(d[0])这里关键参数是编码格式必须指定utf-8否则 Windows 下容易出现UnicodeDecodeError。每一条样本一般包含文本字段、图片字段和标签字段图片字段通常是本地相对路径。不同来源的数据集字段名不一样MVSA 类的数据常见为text、image、label但有的版本叫sentence、img_path。在使用前要统一映射到DataProcess.py里期望的字段名否则会在加载图片时直接报KeyError。文本侧处理的核心是DataProcess.py中定义的 tokenizer 与 padding 逻辑。常见做法是使用transformers库的BertTokenizer并设置max_len128from transformers import BertTokenizer tokenizer BertTokenizer.from_pretrained(bert-base-uncased) encoded tokenizer( the food is great, paddingmax_length, truncationTrue, max_length128, return_tensorspt ) print(encoded[input_ids].shape)这里bert-base-uncased是 BERT 基础版本文本全部转为小写适合英文数据。如果使用中文数据需要换成bert-base-chinese并相应调整词表。max_length128是文本与图片融合时的常见折中太短会截断关键上下文太长会增加 BERT 计算量。项目中的Config.py定义了max_len修改时要注意与DataProcess.py读取处保持一致。2.2 图像侧处理与 ResNet50 的特征提取层图像侧使用 torchvision 的 ResNet50 作为视觉特征提取器。先不急着直接用预训练权重而是看一下DataProcess.py中图片的加载与增强方式from torchvision import transforms from PIL import Image transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) img Image.open(data/sample.jpg).convert(RGB) tensor transform(img) print(tensor.shape) # torch.Size([3, 224, 224])Resize((224, 224))是 ResNet 输入的标准尺寸Normalize采用 ImageNet 统计数据这是从 torchvision 预训练权重中继承来的必要操作。如果不做归一化输入分布和预训练分布不一致早期层输出会偏差较大收敛也会变慢。模型侧加载 ResNet50 时要把最后的全连接层去掉或替换成恒等映射import torchvision.models as models import torch.nn as nn resnet models.resnet50(pretrainedTrue) resnet.fc nn.Identity()关键是这一步fc nn.Identity()使 ResNet 变成纯特征提取器输出的维度是 2048。这个 2048 维向量会与文本侧的 BERT 输出做融合。项目中的每种融合模型都依赖这个 2048 维视觉特征所以在修改 ResNet 结构前要同时检查Models/下所有模型输入维度假设是否一致。2.3 标签与数据集划分的注意事项train.json中的标签一般已编码为 0/1。如果不确定可以用以下命令检查类别分布import json d json.load(open(data/train.json, encodingutf-8)) labels [x[label] for x in d] print(set(labels), len(labels))文本和图片可能存在不匹配的情况比如某条文本是正面但对应图片模糊难辨。这类样本在训练时会造成梯度噪声常见做法是先用一个简单的单模态模型筛查把明显错配的样本剔除或修正。另外要注意Trainer.py中的数据划分它一般会从训练集中再切出验证集划分比例在Config.py的val_ratio字段控制。如果数据量小建议加大验证集比例避免因验证集太小导致准确率波动剧烈。3. 五种融合模型的设计思路与关键代码3.1 简单拼接类NaiveCatModel 与 NaiveCombineModel这两种方法最容易理解也是基线系统。它们的思路是分别用 BERT 得到文本向量用 ResNet50 得到图像向量将两者在特征维度上直接拼接或相加再送入全连接分类器。以NaiveCatModel.py为例核心结构如下import torch.nn as nn from transformers import BertModel import torchvision.models as models class NaiveCatModel(nn.Module): def __init__(self, bert_namebert-base-uncased, num_classes2): super().__init__() self.bert BertModel.from_pretrained(bert_name) self.resnet models.resnet50(pretrainedTrue) self.resnet.fc nn.Identity() # BERT输出768维ResNet输出2048维拼接后2816维 self.classifier nn.Sequential( nn.Linear(768 2048, 512), nn.ReLU(), nn.Dropout(0.3), nn.Linear(512, num_classes) ) def forward(self, input_ids, attention_mask, image): text_feat self.bert(input_ids, attention_maskattention_mask).pooler_output img_feat self.resnet(image) fused torch.cat([text_feat, img_feat], dim-1) return self.classifier(fused)这里pooler_output是 BERT 的 [CLS] 向量经过全连接和 tanh 后的结果维度是 768。直接用pooler_output比取last_hidden_state的均值更常见因为后者需要额外做池化。torch.cat在最后一个维度拼接得到 2816 维。分类器先用 512 维隐藏层做非线性映射再用 Dropout 防止过拟合。NaiveCombineModel与NaiveCatModel的差别在于融合方式Combine 可能是加性融合即text_feat img_feat后进入分类器。加性融合要求两个特征向量的维度一致所以通常会对 2048 维图像特征做一次线性投影到 768 维代码大致为self.img_proj nn.Linear(2048, 768) fused text_feat self.img_proj(img_feat)加性融合比拼接更节省参数但早期层的信息交互较弱。两个模型都适合作为“下限基线”用来证明后面的注意力模型确实有效。3.2 跨模态注意力模型 CMACModelCMACModelCross-Modality Attention Combine Model引入跨模态注意力不再让两个模态的特征独立走到分类器而是让文本特征根据图像特征加权聚合反向同理。这里用最常见的做法把文本特征作为 Query图像特征作为 Key/Value做一次多头注意力。由于 BERT 本身就是 Transformer 结构可以直接复用transformers中的BertAttention不过更可控的手写实现如下import torch import torch.nn as nn import torch.nn.functional as F class CrossAttention(nn.Module): def __init__(self, d_model768, num_heads4): super().__init__() self.num_heads num_heads self.d_model d_model self.d_k d_model // num_heads self.W_q nn.Linear(d_model, d_model) self.W_k nn.Linear(d_model, d_model) self.W_v nn.Linear(d_model, d_model) self.out_proj nn.Linear(d_model, d_model) def forward(self, query, key, value, maskNone): batch_size query.size(0) Q self.W_q(query).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) K self.W_k(key).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) V self.W_v(value).view(batch_size, -1, self.num_heads, self.d_k).transpose(1, 2) scores torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn F.softmax(scores, dim-1) context torch.matmul(attn, V) context context.transpose(1, 2).contiguous().view(batch_size, -1, self.d_model) return self.out_proj(context)d_model768与 BERT 隐藏维度对齐。如果图像特征是 2048 维需要先通过一个 Linear 投影到 768 维才能作为 K 和 V 参与注意力。mask用于遮蔽文本 padding 部分防止无效 token 参与计算。CMACModel 中通常做双向跨模态注意力即文本对图像一次、图像对文本一次然后分别池化后拼接再分类。3.3 输出 Transformer 编码器模型 OTEModelOTEModelOutput Transformer Encoder Model的思路是把两个模态的特征 concat 后再送入一个独立的 Transformer Encoder 层让模态间信息在局部充分交互。实现要点是构造类似 BERT 的 token 序列。文本侧取 BERTlast_hidden_state序列长度 L768 维图像侧取 ResNet 的输出一般池化成一个 2048 维向量再投影成 768 维然后扩展成一维 token与文本 token 拼接成序列长度 L1送入标准 Transformer Encoderfrom transformers import BertConfig from transformers.models.bert.modeling_bert import BertEncoder config BertConfig(hidden_size768, num_hidden_layers2, num_attention_heads8) encoder BertEncoder(config) # text_feat: [batch, L, 768] # img_feat: [batch, 768] - [batch, 1, 768] combined torch.cat([text_feat, img_feat.unsqueeze(1)], dim1) encoded encoder(combined, attention_maskextended_mask)[0] output encoded[:, 0, :] # 取文本 [CLS] 位置这里的extended_mask需要把 padding 位置和图像 token 位置统一考虑。图像 token 是有效位置所以 mask 中对应位置为 1。num_hidden_layers2是常见选择更大的层数会明显增加训练时间但对于小规模数据提升有限。3.4 隐藏状态 Transformer 编码器模型 HSTECModelHSTECModelHidden State Transformer Encoder Combine Model与 OTEModel 的区别在于它不是用 BERT 最后一层的输出作为文本特征而是使用 BERT 每一层的隐藏状态将这些隐藏状态序列与图像特征一起送入 Transformer 编码器进行融合。这种做法的理论依据是BERT 不同层编码不同粒度的语义信息低层偏向词法、高层偏向句法和语义直接只取最后一层会丢失中间层的信息。实现时常用两种方式一是将每层的 pooler 输出拼接成一个序列二是将每层的 [CLS] 向量堆叠成一个张量再作为 Transformer 的输入序列。with torch.no_grad(): outputs bert(input_ids, attention_maskattention_mask, output_hidden_statesTrue) hidden_states outputs.hidden_states # 13个元素第一个是embedding层 # 取第3层、第6层、第9层、第12层的[CLS]向量 cls_list [hidden_states[i][:, 0, :] for i in [3, 6, 9, 12]] multi_layer_feat torch.stack(cls_list, dim1) # [batch, 4, 768]然后把这个[batch, 4, 768]的序列与图像特征序列融合。这样做的好处是模型能自适应地从不同抽象层选取有用信息但代价是计算量增大。如果显存不够可以减少采样层数到两层。3.5 五种模型的对比与选择建议以下表格总结了五种模型的融合方式和适用场景模型融合方式参数量复杂度适用场景NaiveCatModel特征直接拼接低低快速基线NaiveCombineModel投影后加性融合低低验证维度对齐效果CMACModel单向/双向跨模态注意力中中强调文本图像相互引导OTEModelconcat后过Transformer中高高需要序列级交互HSTECModel多层隐藏状态Transformer高最高深层语义与视觉联合建模实际使用中如果数据量小建议先跑 NaiveCatModel 拿到基准。如果 NaiveCatModel 已经过拟合或效果不好再升级到 CMACModel。OTEModel 和 HSTECModel 更适合数据量充足、且需要发论文的场景它们的可解释性也更强可以从注意力权重可视化中分析模态交互。4. 训练流程、验证指标与预测实操4.1 Trainer.py 的训练循环与关键参数Trainer.py承担了训练循环、学习率调度和设备管理的职责。先看最核心的训练步import torch from torch.optim import AdamW from torch.nn import CrossEntropyLoss optimizer AdamW(model.parameters(), lr2e-5, weight_decay0.01) loss_fn CrossEntropyLoss() model.train() for batch in train_loader: input_ids batch[input_ids].to(device) attention_mask batch[attention_mask].to(device) images batch[images].to(device) labels batch[labels].to(device) logits model(input_ids, attention_mask, images) loss loss_fn(logits, labels) optimizer.zero_grad() loss.backward() optimizer.step()这里lr2e-5是从 BERT 微调经验继承来的经典学习率比普通模型的 1e-3 低得多因为 BERT 预训练权重已经很接近局部最优过大的学习率会破坏其学到的语义表示。weight_decay0.01是 AdamW 的常用默认值只对非 bias 和非 LayerNorm 参数生效更佳但项目为了简单直接作用于所有参数问题也不大。Config.py中的关键参数包括batch_size 16 epochs 5 learning_rate 2e-5 warmup_ratio 0.1 print_freq 20batch_size16对于 BERT base 加 ResNet50 的组合在 8GB 显存上运行较稳。如果显存不足可将 batch_size 减半并相应调大梯度累积步数。warmup_ratio0.1表示前 10% 的训练步数内学习率线性上升随后线性衰减这是 Transformer 训练的标准做法。4.2 验证与准确率评估训练过程中Trainer.py会周期性在验证集上计算准确率。核心评估代码逻辑如下model.eval() correct 0 total 0 with torch.no_grad(): for batch in val_loader: logits model(batch[input_ids].to(device), batch[attention_mask].to(device), batch[images].to(device)) preds torch.argmax(logits, dim-1) correct (preds.cpu() batch[labels]).sum().item() total batch[labels].size(0) acc correct / total print(fValidation Acc: {acc:.4f})除了准确率建议额外记录 F1 分数因为情感分析数据常存在类别不平衡。项目中没有显式提供 F1 计算但可以很容易在Trainer.py中加入from sklearn.metrics import f1_score f1 f1_score(batch[labels].cpu(), preds.cpu(), averagebinary)二分类场景用averagebinary即可。如果是多分类改成averagemacro。验证时务必调用model.eval()它会关闭 Dropout 和 BatchNorm 的 training 状态否则推理结果会有随机性。4.3 用训练好的模型跑 test_without_label.txt项目提供的test_without_label.txt是没有标签的测试集用于生成提交结果。加载模型后输出预测结果并写入文件import torch from Models.NaiveCatModel import NaiveCatModel model NaiveCatModel() state_dict torch.load(checkpoints/best_model.pt, map_locationcpu) model.load_state_dict(state_dict) model.eval() predictions [] with torch.no_grad(): for sample in test_loader: logits model(sample[input_ids], sample[attention_mask], sample[images]) pred torch.argmax(logits, dim-1).cpu().tolist() predictions.extend(pred) with open(prediction.txt, w, encodingutf-8) as f: for p in predictions: f.write(str(p) \n)map_locationcpu能避免在没有 GPU 的环境下加载模型报错。best_model.pt是Trainer.py在验证集上表现最好的 checkpoints 文件保存路径在Config.py的checkpoint_dir中定义。加载时如果报Missing key(s) in state_dict多半是保存时用了DataParallel对应处理办法是把键名前缀去掉state_dict torch.load(checkpoints/best_model.pt) from collections import OrderedDict new_state_dict OrderedDict() for k, v in state_dict.items(): new_state_dict[k.replace(module., )] v model.load_state_dict(new_state_dict)4.4 显存不足与训练速度优化BERT base 和 ResNet50 同时前向反向显存占用大约 4GB 到 7GB具体取决于 batch size 和是否启用中间层输出。显存不足时采用以下优先策略第一降低 batch size 并增加梯度累积accumulation_steps 4 optimizer.zero_grad() for i, batch in enumerate(train_loader): loss loss_fn(model(...), labels) / accumulation_steps loss.backward() if (i 1) % accumulation_steps 0: optimizer.step() optimizer.zero_grad()第二给图像分支加torch.no_grad()前提是只微调 BERT 不微调 ResNet。由于多数多模态情感分析中文本信息量更大这种做法能显著减少显存with torch.no_grad(): img_feat resnet(images)第三使用混合精度训练。torch 1.8 中可借助torch.cuda.ampfrom torch.cuda.amp import autocast, GradScaler scaler GradScaler() with autocast(): logits model(input_ids, attention_mask, images) loss loss_fn(logits, labels) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()autocast只在 GPU 上生效CPU 环境会自动忽略。5. API 封装与 requirements 环境安装5.1 requirements.txt 版本约束说明项目 requirements.txt 给出了固定版本chardet4.0.0 numpy1.22.2 Pillow9.2.0 scikit_learn1.1.1 torch1.8.2 torchvision0.9.2 tqdm4.63.0 transformers4.18.0其中torch1.8.2和transformers4.18.0的组合是经过验证的稳定搭配。torchvision0.9.2必须与 torch 1.8.2 配套否则可能出现算子不兼容或 ResNet 权重加载失败。安装命令pip install -r requirements.txt建议在创建虚拟环境后安装python -m venv venv source venv/bin/activate # Windows 下为 venv\Scripts\activate pip install --upgrade pip pip install -r requirements.txt如果你的 Python 版本是 3.9 或更低直接安装即可。如果 Python 是 3.10torch 1.8.2 的预编译包可能无法安装这时需要改用更高版本的 torch比如torch1.12.1和torchvision0.13.1但注意 transformers 4.18.0 在高版本 torch 下仍可运行。5.2 通过 main.py 启动完整流程main.py是整个项目的入口一般支持三个模式训练、验证、测试。命令格式通常为python main.py --mode train --model NaiveCat python main.py --mode eval --model CMAC --checkpoint checkpoints/best_model.pt python main.py --mode predict --model OTE --checkpoint checkpoints/best_model.pt --input data/test_without_label.txt--model参数值对应 Models 目录下的类名注意NaiveCat对应NaiveCatModel.py中的NaiveCatModel。如果运行时报ModuleNotFoundError多半是因为没有在main.py中把Models目录加入 Python 路径import sys import os sys.path.append(os.path.join(os.path.dirname(__file__), Models))在main.py中训练模式的白话流程是读取 Config 中的参数实例化数据加载器创建模型调用 Trainer 的train()方法。如果想快速验证环境是否正常可以先用--epochs 1 --batch_size 4跑一个小样本测试python main.py --mode train --model NaiveCat --epochs 1 --batch_size 4这里--epochs 1只跑一个 epoch目的不是得到好模型而是确认数据管道、前向传播和反向传播都没有报错。5.3 使用 transformers 的 pipeline 做快速调试在训练完整模型前可以用 transformers 的pipeline快速验证文本侧 BERT 的情感判断准确性作为对比参照from transformers import pipeline cls pipeline(sentiment-analysis, modelnlptown/bert-base-multilingual-uncased-sentiment) print(cls(The food was amazing but the service was slow.))这个多语言模型输出 1 到 5 星的评价可以把它映射到二分类上对比一下。注意这里只是为了快速感知文本单模态的上限并不是项目最终的验证方式。多模态融合模型的目标是超过文本单模态准确率如果在一组数据上融合模型反而不如纯文本说明图片模态带来了噪声需要检查图像预处理或数据配对质量。6. 融合效果的可解释性分析与一个实用技巧多模态模型的劣势在于难以解释模型到底依据什么做出判断。对情感分析而言一个直接验证方法是对跨模态注意力矩阵做可视化观察文本中的情感词是否与图像显著区域形成对应。以 CMACModel 为例在 forward 中把注意力权重保存下来class CMACModelWithAttn(nn.Module): def __init__(self): super().__init__() self.attn_weights None def forward(self, input_ids, attention_mask, image): # ... 前面相同 scores torch.matmul(Q, K.transpose(-2, -1)) / (self.d_k ** 0.5) self.attn_weights F.softmax(scores, dim-1) # ... 后续操作然后用热力图展示图像侧的注意力响应。对视觉分支需要把 ResNet 输出的特征图而不是最终池化向量作为注意力目标这要求修改 ResNet 的 forward。最省力的方式是使用torchvision.models.feature_extraction创建特征提取器from torchvision.models.feature_extraction import create_feature_extractor model models.resnet50(pretrainedTrue) feature_extractor create_feature_extractor( model, return_nodes{layer4.2.relu: feature_map} ) out feature_extractor(image_tensor)[feature_map]返回的feature_map形状是[batch, 2048, 7, 7]每个空间位置对应原图 32x32 的感受野。把特征图按通道维求平均再 resize 到原图尺寸就能叠加可视化import torch.nn.functional as F import numpy as np from PIL import Image feat out[0].mean(dim0) # [7, 7] feat F.interpolate(feat.unsqueeze(0).unsqueeze(0), size(224, 224), modebilinear) heatmap (feat.squeeze().numpy() - feat.min().item()) / (feat.max().item() - feat.min().item()) img np.array(Image.open(data/sample.jpg).convert(RGB)) overlay (img * 0.5 (heatmap[..., None] * 255) * 0.5).astype(np.uint8)这种做法在论文配图和 debug 中很实用如果一张图片的情感是正面的但热力图高亮区域集中在无关背景上说明模型学到的是背景与标签的虚假相关性此时应该检查数据集或增加正则化。最后分享一个实验技巧固定一个统一的特征提取 backbone保持 BERT 和 ResNet 的预训练权重不更新只微调融合层和分类器先用这个方案跑通五种模型。这样可以快速对比不同融合方法的相对优劣排除主干网络微调带来的干扰。等选定最佳融合模型后再解锁 BERT 或 ResNet 的微调通常能再稳定提升 1 到 3 个百分点。训练时记得在Config.py中关闭与当前模式不匹配的日志输出避免每步都打印导致print_freq形同虚设。本文还有配套的精品资源点击获取
返回列表