简介:一套基于PyTorch实现的多模态虚假新闻检测系统源码,面向深度学习研究者、NLP方向学生及舆情分析开发者。方案融合BERT预训练模型与ResNet卷积神经网络,分别提取文本深层语义与图像视觉特征,并在微博谣言数据集上完成训练与评估,同时引入对比学习机制强化真实新闻与虚假新闻的判别能力,适用于社交媒体虚假信息识别等任务。压缩包共收纳21个文件,整体仅1.48MB,以12个Python脚本为主体,涵盖模型结构、训练流程、数据预处理与工具函数;另有4个文本说明、3个CSV数据及docx、md格式的配套资料,方便对照阅读与二次开发。目前已有96人浏览/学习。从实战角度看,整套代码不仅提供完整的模型结构与可运行训练入口,还通过说明文档梳理了参数配置、数据组织方式与扩展思路,适合作为多模态分类项目的起步模板;通过阅读源码可学习BERT文本特征与ResNet图像特征的融合方式,以及对比学习如何优化表征并复现微博谣言数据上的实验结果。
1. 多模态虚假新闻检测为什么值得复现:先立住技术选型,再谈落地难度
拿到一个“基于PyTorch的多模态虚假新闻检测系统”代码包,很多人的第一反应是先把环境配好再跑通训练脚本。但这个方向真正值得花时间的点不在训练本身,而在“文本和图像特征怎么对齐”和“对比学习到底给分类任务带来了什么”。这条技术路线在PyTorch生态里已经非常成熟:BERT负责文本语义,ResNet负责图像内容,两者各出一个特征向量,再通过对比学习让同一篇博文的图文表征靠得更近。它作为论文复现、毕业设计或企业舆情系统的原型,性价比都很高。适合对PyTorch有基础、想把多模态模型从概念落到代码的工程师,也适合想拿公开数据集做消融实验的研究者。这篇笔记按“数据处理 -> 模型搭建 -> 对比学习训练 -> 排障 -> 评估验证”的顺序展开,所有代码都按可运行的标准来写。
2. 微博谣言数据集预处理:从原始标注到BERT与ResNet都能吃的样子
2.1 数据集字段与标注含义:先看清JSON里有什么
公开的微博不实信息数据集一般以JSON或CSV形式发布,每条样本包含“微博文本内容”“图片URL列表”“标注(谣言/非谣言)”“发布时间”“转发数”“点赞数”等字段。做多模态检测时,真正用到的只有文本、图片和标签三样,其余字段可以作为辅助特征,但前期建议先不放,避免引入过多噪声。
不同版本的微博数据集在字段命名上不太统一,有的用weibo_text,有的用text,图片字段有的是image_url,有的是pic_list。拿到数据的第一件事不是写训练代码,而是写一个探查脚本,把字段名、类型、缺失情况一次性打印出来。
import json from collections import Counter with open("weibo_dataset.json", "r", encoding="utf-8") as f: data = json.load(f) print("样本总数:", len(data)) print("字段名:", list(data[0].keys())) print("标签分布:", Counter([item["label"] for item in data])) # 检查文本和图片字段的缺失率 missing_text = sum(1 for item in data if not item.get("text")) missing_img = sum(1 for item in data if not item.get("image_url")) print("缺失文本:", missing_text, "缺失图片:", missing_img)这段代码做了三件事:确认数据量、确认字段名、确认标签分布。标签分布这一步很重要,微博谣言数据集的谣言与非谣言比例通常不是1:1,有的版本甚至接近1:3,这个比例会直接影响后面评估指标的选择。缺失文本和缺失图片的数量也要先摸清,如果缺失比例超过10%,处理策略要相应调整。
2.2 文本清洗与BERT Tokenizer处理:统一到128个token以内
BERT的输入是token序列,但原始微博文本里有很多对语义没帮助的内容:HTML转义符、@用户昵称、URL、表情符号的文本表示。这些东西如果不清理,tokenizer会把一个长URL拆成几十个碎token,白白占用序列长度,BERT根本看不完正文内容。
常见的清洗顺序是:先去HTML实体,再替换URL和@用户,最后保留中文字符和基本标点。清洗做完后,用HuggingFace的BertTokenizer把文本转成input_ids、attention_mask、token_type_ids三件套。
from transformers import BertTokenizer tokenizer = BertTokenizer.from_pretrained("bert-base-chinese") MAX_LEN = 128 def clean_text(text): import re text = re.sub(r"<[^>]+>", "", text) # 去HTML标签 text = re.sub(r"https?://\S+", "【网页链接】", text) # URL替换为占位符 text = re.sub(r"@[\w\-\u4e00-\u9fa5]+", "【用户】", text) # @用户替换 text = re.sub(r"\s+", " ", text).strip() return text def encode_text(text, max_len=MAX_LEN): cleaned = clean_text(text) encoded = tokenizer( cleaned, max_length=max_len, padding="max_length", truncation=True, return_tensors="pt" ) return { "input_ids": encoded["input_ids"].squeeze(), "attention_mask": encoded["attention_mask"].squeeze() }这段代码里的关键参数是max_length=128。微博文本大多比较短,128个token足够覆盖绝大多数内容,更重要的是,序列长度直接决定BERT前向推理的显存占用,从512降到128,显存能省下近一半。return_tensors="pt"让tokenizer直接返回PyTorch张量,省去手动转换。注意这里没有返回token_type_ids,因为单文本分类任务用不到,返回了也只能增加数据体积。
2.3 图片解码与缩放:ResNet不喜欢原图分辨率
ResNet的输入要求是固定尺寸的RGB张量,torchvision的标准是224×224。但我们从网上下载的微博图片分辨率参差不齐,有的大图5000×3000,有的小图只有80×80,直接把原图送入网络会导致两个问题:一是数据加载变慢,二是不同样本经过ResNet的池化层后输出尺寸不一致,无法对齐特征。
图片处理的基本流程是:下载或读取图像文件后用PIL解码,转换成RGB模式,按比例缩放使短边到256,再中心裁剪出224×224。这个“先缩放再裁剪”的顺序比直接拉伸更保真,不会让图像变形。
from PIL import Image from torchvision import transforms img_transform = transforms.Compose([ transforms.Resize((256, 256)), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) ]) def load_image(img_path): try: img = Image.open(img_path).convert("RGB") return img_transform(img) except Exception: # 图片损坏或下载失败时返回全零张量,后续由Dataset层处理 return None这里有两个参数值得注意:Resize((256, 256))的256是经验值,比目标尺寸224稍大一点,给随机裁剪留出空间;Normalize用的均值方差是ImageNet的统计值,因为我们要加载在ImageNet上预训练的ResNet权重,输入分布与预训练分布保持一致才能发挥迁移学习的效果。如果直接使用没有预训练的随机初始化ResNet,归一化参数用默认的0.5也能跑,但效果会明显变差。
2.4 Dataset封装与训练集划分:把数据读取从训练循环里解放出来
数据清洗和编码做完后,要把所有逻辑封装进torch.utils.data.Dataset子类。多模态数据集的难点在于文本和图像要成对返回,而且图像存在下载失败、文件损坏的可能,这部分必须在Dataset里做容错处理,不能在训练循环里临时判断。
import torch from torch.utils.data import Dataset class MultimodalDataset(Dataset): def __init__(self, data, tokenizer, img_root, max_len=128): self.samples = [] for item in data: text = encode_text(item["text"], max_len)["input_ids"] mask = encode_text(item["text"], max_len)["attention_mask"] img_path = f"{img_root}/{item['img_id']}.jpg" self.samples.append({ "input_ids": text, "attention_mask": mask, "img_path": img_path, "label": torch.tensor(item["label"], dtype=torch.long) }) def __len__(self): return len(self.samples) def __getitem__(self, idx): sample = self.samples[idx] img = load_image(sample["img_path"]) if img is None: img = torch.zeros(3, 224, 224) # 坏图补零,而不是跳过样本 return { "input_ids": sample["input_ids"], "attention_mask": sample["attention_mask"], "image": img, "label": sample["label"] }这个Dataset有两个设计细节。第一,文本编码在__init__里提前算好,训练时直接查表取用,避免每个epoch都重复做tokenize;第二,图像加载失败时返回全零张量而不是丢弃该样本,因为如果测试集里有一张坏图就跳过,会导致评估集不完整,指标对比失去意义。全零张量会让ResNet输出一个无意义特征,但配合后面的对比学习,模型会学会对此类样本不可靠,实测中比直接跳过更安全。
3. 搭建双塔特征提取网络:BERT与ResNet的接入方式与输出对齐
3.1 文本塔:BERT的CLS向量与最后一层Hidden State怎么取舍
BERT作为文本编码器,最常见的接入方式是取[CLS]位置的输出向量作为整条文本的表示。这个向量经过12层Transformer的交互之后,理论上聚合了全局语义信息。另一种做法是取最后一层所有token的向量做平均池化或最大池化,在某些短文本任务上平均池化比CLS表现更稳定。
在微博谣言检测这个任务上,我一般保留CLS向量,理由是对话立场和语气等信号在CLS里编码更充分,而且CLS向量是后续对比学习投影头最自然的输入。实现上用transformers库直接构建模型,不额外写BERT结构,省心也稳定。
from transformers import BertModel class TextEncoder(torch.nn.Module): def __init__(self, pretrained="bert-base-chinese", freeze=False): super().__init__() self.bert = BertModel.from_pretrained(pretrained) self.feat_dim = self.bert.config.hidden_size # 768 if freeze: for param in self.bert.parameters(): param.requires_grad = False def forward(self, input_ids, attention_mask): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) return outputs.last_hidden_state[:, 0, :] # [batch, 768]outputs.last_hidden_state[:, 0, :]取的就是每个序列的第一个token,即CLS位置的向量。self.bert.config.hidden_size是768,这个数字后面做特征对齐时要用到,建议写成从config读取而不是硬编码。关于是否冻结BERT参数,下面搭配整体训练策略时再细说。
3.2 图像塔:去掉ResNet分类头的特征提取写法
torchvision里的ResNet模型默认带着一个1000类的全连接分类头,用作特征提取时要把fc层去掉或替换成恒等映射。ResNet18的最后一个卷积块输出512维特征图,经过全局平均池化得到512维向量;ResNet50则是2048维。选哪个主要看显存和精度要求,ResNet18在微博图片这类场景下已经够用,显存占用也友好。
import torchvision.models as tv_models class ImageEncoder(torch.nn.Module): def __init__(self, base="resnet18", freeze=False): super().__init__() if base == "resnet18": self.backbone = tv_models.resnet18(weights=tv_models.ResNet18_Weights.IMAGENET1K_V1) self.feat_dim = 512 else: self.backbone = tv_models.resnet50(weights=tv_models.ResNet50_Weights.IMAGENET1K_V2) self.feat_dim = 2048 # 去掉最后的分类层,保留特征提取部分 self.backbone.fc = torch.nn.Identity() if freeze: for param in self.backbone.parameters(): param.requires_grad = False def forward(self, image): return self.backbone(image) # [batch, feat_dim]把fc换成Identity()是最省事的做法,ResNet前向计算到全局平均池化后直接输出特征向量,不会多算一个1000维的线性层。这里用了torchvision.models的预训练权重枚举,比旧版传字符串的方式可读性更好。weights参数在PyTorch 2.x里是推荐写法,传IMAGENET1K_V1或IMAGENET1K_V2都能自动下载权重。
3.3 特征对齐与融合:768和512怎么拼到一起
双塔输出的维度不一致,这是必须面对的第一个问题。融合方案有几种,最简单的concat,除此之外还有加权求和、cross-attention等方式。在微博谣言检测这个任务规模上,concat到一个维度后接全连接层是最稳妥的做法,信息损失最小,也方便后面做消融实验对比单模态基线。
class MultimodalModel(torch.nn.Module): def __init__(self, text_encoder, image_encoder, num_classes=2, fusion_dim=256, use_projection=True): super().__init__() self.text_encoder = text_encoder self.image_encoder = image_encoder self.fusion_dim = fusion_dim self.use_projection = use_projection total_dim = text_encoder.feat_dim + image_encoder.feat_dim self.classifier = torch.nn.Sequential( torch.nn.Linear(total_dim, fusion_dim), torch.nn.ReLU(), torch.nn.Dropout(0.3), torch.nn.Linear(fusion_dim, num_classes) ) # 对比学习用的投影头,把融合特征映射到低维对比空间 self.projection_head = torch.nn.Sequential( torch.nn.Linear(total_dim, 128), torch.nn.ReLU(), torch.nn.Linear(128, 128) ) def forward(self, input_ids, attention_mask, image, return_feature=False): text_feat = self.text_encoder(input_ids, attention_mask) img_feat = self.image_encoder(image) fused_feat = torch.cat([text_feat, img_feat], dim=-1) logits = self.classifier(fused_feat) if return_feature: proj_feat = self.projection_head(fused_feat) return logits, proj_feat return logits这里的核心参数是fusion_dim=256,既不过大导致过拟合,也不太小导致信息瓶颈。Dropout(0.3)是针对融合层的经验值,多模态场景下两个特征来源差异大,融合层比单模态分类器更容易过拟合。projection_head输出128维向量,这是给对比学习专门预留的特征空间,与分类分支共享底层特征提取器,但各用各的输出头。
3.4 冻结策略:BERT和ResNet到底要不要一起微调
一个完整的多模态模型参数量很大:BERT base大约1.1亿,ResNet18大约1200万,全量微调对显存和训练时间都是考验。常见做法分三种:全部冻结只训练融合层、部分冻结微调深层、全部微调。在微博数据集规模(几万到十几万样本)下,我一般选择冻结BERT的前8层,微调后4层和ResNet的最后一个stage,这个策略在公开数据集上表现接近全量微调,但显存占用降低约40%。
def set_partial_freeze(model, freeze_text_layers=8): # 冻结BERT前8层 for name, param in model.text_encoder.bert.named_parameters(): layer_idx = int(name.split(".")[2]) if "layer" in name else -1 if 0 <= layer_idx < freeze_text_layers: param.requires_grad = False层编号解析依赖transformers的内部命名规则,bert.encoder.layer.0.attention.self.query.weight这种格式里split(".")[2]取到的就是层序号。冻结浅层非常关键,BERT浅层编码的是词法和句法特征,通用性强,深层才编码任务相关语义,微博这种短文本场景尤其如此。ResNet则建议冻结前三个stage,只微调最后一个stage和全局池化前的卷积块。
4. 用对比学习拉近同源样本:训练策略与关键参数设置
4.1 对比学习在这个任务里解决什么问题
单纯用交叉熵训练多模态分类器,模型可能学到“文本说转发、图片是风景”这类表面关联,鲁棒性不足。对比学习的核心目标是拉近同一个样本图文表征的距离,推远不同样本的图文表征,让模型学会“这条微博的图和文在说的是同一件事”这个语义对齐关系。
具体到实现,每个训练样本天然的图文对就是正样本对,同一个batch里其他样本的图文特征就是负样本。这样不需要额外构造数据,也不依赖标注信息,直接利用多模态数据的配对结构。这比用同一条文本在不同数据增强下的视图做正对更自然。
import torch.nn.functional as F def info_nce_loss(features, temperature=0.07): """features: [batch, 2, proj_dim],0是文本特征,1是图像特征""" text_feat = features[:, 0, :] img_feat = features[:, 1, :] text_feat = F.normalize(text_feat, dim=-1) img_feat = F.normalize(img_feat, dim=-1) logits = torch.matmul(text_feat, img_feat.T) / temperature batch_size = text_feat.size(0) labels = torch.arange(batch_size, device=features.device) loss = F.cross_entropy(logits, labels) return losstemperature=0.07是对比学习里常用的起点值,它控制相似度分布的尖锐程度。温度太低会让训练不稳定,温度太高会让正负样本的区分度变弱。labels用的是对角线索引,因为第i个样本的图文特征正确配对时,text_feat[i]应该和img_feat[i]最相似,对角线就是监督信号。
4.2 图文联合训练:主分类损失和对比损失怎么加权
对比学习只是辅助约束,不能让它的 loss 压过分类 loss 太多,否则模型会把精力全放在对齐图文上,反而忽略了分类任务。实践上一般让分类损失做主导,对比损失作为正则项。这里的关键参数就是权重系数 lambda,从0.1开始调,如果分类指标提升不明显就往上加,如果训练震荡就往下减。
def train_step(model, batch, optimizer, lambda_cl=0.1, temperature=0.07): input_ids = batch["input_ids"].to(device) attention_mask = batch["attention_mask"].to(device) image = batch["image"].to(device) labels = batch["label"].to(device) logits, proj_feat = model(input_ids, attention_mask, image, return_feature=True) cls_loss = F.cross_entropy(logits, labels) # 构造对比学习输入:让投影头分别处理文本和图像特征 text_feat = model.text_encoder(input_ids, attention_mask) img_feat = model.image_encoder(image) fused_text = torch.cat([text_feat, torch.zeros_like(img_feat)], dim=-1) fused_img = torch.cat([torch.zeros_like(text_feat), img_feat], dim=-1) proj_text = model.projection_head(fused_text) proj_img = model.projection_head(fused_img) features = torch.stack([proj_text, proj_img], dim=1) cl_loss = info_nce_loss(features, temperature) total_loss = cls_loss + lambda_cl * cl_loss optimizer.zero_grad() total_loss.backward() optimizer.step() return total_loss.item(), cls_loss.item(), cl_loss.item()这段代码里对比学习分支的做法是把文本特征和图像特征分别拼进完整特征向量,再送入投影头。理由是投影头接收的是融合特征的维度,不能直接对768维文本特征和512维图像特征分别投影,否则对比空间的维度不一致。用零填充补位虽然有点浪费计算,但保证了投影头输入的语义空间与分类分支一致。
4.3 训练循环与学习率设置:BERT用低学习率,新层用高学习率
多模态模型的参数更新频率天然不一致,BERT层如果和分类层用同一个学习率,BERT会被冲乱。常见做法是设置分组学习率:BERT和ResNet的backbone用小学习率(2e-5到5e-5),融合层、分类头和投影头用大学习率(1e-4到3e-4)。优化器用AdamW,这是PyTorch生态里预训练模型微调的事实标准。
from transformers import AdamW def build_optimizer(model, lr_backbone=3e-5, lr_head=1e-4): backbone_params = [] head_params = [] for name, param in model.named_parameters(): if not param.requires_grad: continue if "classifier" in name or "projection_head" in name: head_params.append(param) else: backbone_params.append(param) return AdamW([ {"params": backbone_params, "lr": lr_backbone}, {"params": head_params, "lr": lr_head} ])lr_backbone=3e-5是BERT类模型的标准起点,1e-4对新初始化的层非常合适。训练轮数上,微博数据集几万样本的规模,5到8个epoch基本收敛,再多就会过拟合。训练过程中每半个epoch记录一次验证集准确率,如果连续两个epoch没有提升,就把学习率降到当前值的0.1倍,这是最省心的调度策略。
5. 多模态训练避坑指南:从显存OOM到评估指标虚高
5.1 显存溢出:OOM的锅不只在batch size
现象:训练刚开始就报CUDA out of memory,把batch size从32降到8还是崩。原因:BERT的序列长度、ResNet输入分辨率、数据加载的pin_memory设置、多个张量同时驻留显存,这些因素叠加在一起,单靠调batch size收效甚微。解决:先看总显存占用结构,用torch.cuda.memory_summary()确认哪些中间变量占空间。最常见的优化是把MAX_LEN从128降到96,把 image 的分辨率保持224不变,同时给 DataLoader 加pin_memory=True但把num_workers调低到2或4。如果还爆,启用梯度累积,每两个batch更新一次参数。
提示:微博文本平均长度只有30到40个token,96的序列长度极少截断正文,优先压缩文本序列而不是图片分辨率。
5.2 BERT权重下载失败与参数路径混乱
现象:本地已经下载过bert-base-chinese,但程序每次启动都从HuggingFace Hub重新下载,网络不稳时直接报连接错误。原因:from_pretrained默认检查的是缓存目录,如果之前下载中断或缓存目录环境变量没配对,就会重复下载。解决:手动指定本地路径加载,提前用snapshot_download把模型拉到指定目录,或者从HuggingFace的模型页手动下载配置文件、tokenizer文件和权重文件,放到项目下的./weights/bert-base-chinese/里:
from transformers import BertModel, BertTokenizer local_bert_path = "./weights/bert-base-chinese" tokenizer = BertTokenizer.from_pretrained(local_bert_path) text_encoder = BertModel.from_pretrained(local_bert_path)这段代码的关键在“一次性指定本地路径”。BERT相关的文件有5个:config.json、pytorch_model.bin、vocab.txt、tokenizer_config.json、special_tokens_map.json,缺任何一个都会加载失败。另一个隐蔽的坑是文件名不一致,比如权重叫pytorch_model.bin但代码里期望model.bin,所以下载后先检查文件名再写路径。
5.3 ResNet预训练权重与PyTorch版本不匹配
现象:加载ResNet时提示Parameter 'fc.weight' of size (1000, 2048) doesn't match the size (2048, 2048)之类的错误。原因:旧代码里用models.resnet50(pretrained=True),新版本torchvision已经改为weights参数,而本地的解决办法是用不符合版本的下载方式。解决:明确用枚举指定权重版本,不要用布尔值。如果本地已有历史权重文件,可以用torch.load手动载入后过滤掉fc层对应的键值再load_state_dict。
import torchvision.models as tv_models state_dict = torch.load("./weights/resnet50_imagenet.pth") state_dict.pop("fc.weight", None) state_dict.pop("fc.bias", None) backbone = tv_models.resnet50(weights=None) backbone.fc = torch.nn.Identity() backbone.load_state_dict(state_dict, strict=False)strict=False是这里的关键参数,它允许状态字典缺失或多余部分不会报错。如果直接strict=True,就会因为fc层被移除而报key不匹配。注意weights=None表示不加载任何预训练权重,完全由手动方式注入,这也绕开了PyTorch新版本对下载来源的校验。
5.4 微博数据集的图片缺失与字段不一致
现象:训练进度到某个epoch突然变慢,或者准确率一直在50%上下跳动。排查后发现数据里图像的img_id和实际文件名对不上,有的图片在下载过程中损坏,PIL直接抛异常。原因:公开数据集的图片URL经常失效,尤其那些多年以前的微博配图,源站已经删除或防盗链。解决:除了在Dataset里对坏图返回全零张量之外,还要在建Dataset之前做一次全量图片存在性检查,统计坏图比例。如果坏图比例超过5%,建议把对应样本的文本塔输出用作主特征,对图像分支做置零处理而不是直接剔除样本,保持训练集和测试集的分布一致。
5.5 对比学习不收敛或拉低分类指标
现象:加了对比学习之后分类准确率反而下降,或者对比损失降不下去,一直维持在0.5以上。原因:最可能是温度参数和lambda权重不匹配。温度太低时负样本相似度梯度几乎为零,投影头学不动;lambda太大时模型只顾着对齐图文对,忽略分类标注。另一个隐蔽问题是对比学习里的正负样本构造错误,比如把同一个样本的文本与文本当成正对,那就没有意义了。解决:先固定分类损失单独训练3个epoch,再开启对比学习分支。温度从0.07开始,lambda从0.05开始,观察对比损失如果10个step内没有下降趋势,把温度调大一点,温度调大仍然不降就检查特征是否做了L2归一化。
6. 评估与验证:用消融实验和特征可视化验证系统真的有效
6.1 评估指标:宏平均F1比准确率更可靠
微博谣言数据集中非谣言样本通常多于谣言样本,准确率会被多数类抬高。一个模型把所有样本都判为非谣言,准确率也能到70%以上,这在业务场景里完全不可用。评估时用宏平均F1(macro F1)和混淆矩阵,按谣言类别单独看召回率。实现上直接用sklearn的f1_score(average='macro')和confusion_matrix,不需要绕弯。
from sklearn.metrics import f1_score, classification_report preds, gt = [], [] model.eval() with torch.no_grad(): for batch in val_loader: logits = model(batch["input_ids"], batch["attention_mask"], batch["image"]) preds.extend(torch.argmax(logits, dim=-1).cpu().tolist()) gt.extend(batch["label"].cpu().tolist()) print("macro F1:", f1_score(gt, preds, average="macro")) print(classification_report(gt, preds, target_names=["非谣言", "谣言"]))评估的时候模型要切到eval()模式并关闭梯度计算,这两步能明显降低显存占用和推理时间。分类报告里的每一行都要看,特别是谣言类的precision和recall,如果recall偏低说明模型漏掉了很多谣言样本,这是生产场景里最不能接受的失误。
6.2 用特征可视化验证对比学习是否真的生效
训练完的模型不能只看指标就交差,还要验证对比学习有没有真的把图文特征拉到一起。做法是把测试集里的图文对分别过模型编码器取投影特征,降维后用颜色标注正负样本,观察图文特征是否聚类。
from sklearn.manifold import TSNE import matplotlib.pyplot as plt text_feats, img_feats, labels = [], [], [] model.eval() with torch.no_grad(): for batch in val_loader: tf = model.text_encoder(batch["input_ids"], batch["attention_mask"]) imf = model.image_encoder(batch["image"]) text_feats.append(tf) img_feats.append(imf) labels.extend(batch["label"].tolist()) text_feats = torch.cat(text_feats).cpu().numpy() img_feats = torch.cat(img_feats).cpu().numpy() tsne = TSNE(n_components=2, perplexity=30, random_state=42) all_feats = tsne.fit_transform(np.concatenate([text_feats[:500], img_feats[:500]]))TSNE(n_components=2, perplexity=30)是可视化常用配置,perplexity太大或太小都会让分布变得松散。抽样500个样本是经验值,样本太多TSNE的计算时间会急剧增加。如果把对应图文对的点距离画出来,正样本对的距离比负样本对小,说明对比学习确实起作用了,这时候模型的泛化能力才有保障。
6.3 消融实验的最小配置
要证明多模态和对比学习的价值,需要跑三组消融:只用BERT文本、只用ResNet图像、双塔不加对比学习、双塔加对比学习。这四组配置可以在同一个模型框架里通过开关切换,每组在相同的数据划分和随机种子下跑,记录macro F1和谣言类recall。做消融时最忌讳的是各跑各的随机种子,最后指标差异分不清是模型贡献还是数据波动。固定随机种子是必须的:
def set_seed(seed=42): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) import numpy as np, random np.random.seed(seed) random.seed(seed)6.4 从复现到改进:值得尝试的进阶方向
模型跑通之后,如果想往业务或论文方向走,有几个性价比高的改进点。第一个是把ResNet18换成ResNet50并对比效果,确认是不是图像特征容量拖了后腿。第二个是在融合层加一层简单的cross-attention,让文本特征和图像特征在融合前互相加权。第三个是用torch.onnx.export把训练好的模型导出为ONNX格式做推理加速验证,这对上线部署很有参考意义。我在做完对比学习验证之后,习惯把所有实验的指标、参数、随机种子记录在一个表格里再决定下一步动哪里,项目里这个“三塔配置”在微博数据集上最稳定的是双塔加对比学习,希望帮到你。
本文还有配套的精品资源,点击获取