news 2026/9/11 19:09:19

三种文本分类模型对比与Python复现:TextGCN、TextING与LEAM解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
三种文本分类模型对比与Python复现:TextGCN、TextING与LEAM解析

简介:面向自然语言处理课程设计与期末项目,这份资源复现了TextGCN、TextING、LEAM三种经典文本分类方法,提供完整Python源码与详细注释。TextGCN基于图卷积网络,TextING基于图神经网络,LEAM则利用标签嵌入注意力机制,三个模型均包含数据预处理、构图、训练与可视化等完整流程,注释对关键函数和参数作了说明,适合计算机、人工智能等专业学生作为课设作业参考,也可供初学者对照学习模型细节。压缩包共91个文件,包含32个Python脚本、10个PDF说明、多个ipynb笔记及npz/npy数据文件,总大小约805MB,目录按三种模型分别组织,并附带README与环境说明,便于快速运行。资源内代码已测试运行成功,可直接修改扩展,也可用于课程设计、毕业设计或项目初期演示。已有437人学习下载,适合需要理解文本分类原理并动手复现的同学。

1. 三种文本分类方法在同一份期末数据上,凭什么能差出十个点

做 NLP 期末大作业,最常遇到的不是模型跑不起来,而是三个模型都能跑,效果却对不上论文。TextGCN、TextING、LEAM 是三种极具代表性的文本分类方法:TextGCN 把整份语料建成一张文档-词大图,走半监督传导式学习;TextING 把每篇文档单独建词图,新样本来了直接出结果;LEAM 完全不建图,用标签词向量引导文本里的注意力。同一份数据上,这三套思路常常能拉开 5~10 个百分点的准确率差,差距不在调参,而在建模假设本身。这篇文章按数据准备、三个模型逐一复现、对比与集成展开,每个模型都给出带注释的 Python 源码、关键参数和踩坑记录,适合 NLP 课设与期末大作业复现,也适合想快速横向验证这三条技术路线的工程师。所有代码基于 Python 3.10,依赖越少越好,能跑通的地方全部给出了最小命令。

2. 复现前的数据准备与 Python 环境配置:先固定词表和词向量

2.1 文本分类数据集怎么选:先把规模限死在够用的范围

TextGCN 的构图开销随词表和文档数线性增长,三套模型共用一份数据时,最省心的选择是 20 个类别以内、1 万到 3 万条的短文本。中文场景我用的是 THUCNews 的 10 类子集,英文可以直接用 sklearn 内置的 20Newsgroups。无论选哪个,先把目录固定成下面三行格式,后面所有代码都按这个格式读:

data/train.txt 标签\t文本 data/dev.txt 标签\t文本 data/test.txt 标签\t文本

这里有一个容易被略过的细节:dev 集不能省。TextGCN 的 epoch 数、TextING 的消息传递步数、LEAM 的 temperature,三个模型都要靠 dev 来挑 checkpoint。没有 dev 时只能按训练 loss 存模型,结果往往是训练集准确率很漂亮,测试集直接掉下去。

2.2 分词、固定词表与 OOV 处理的 Python 代码

中文分词直接用 jieba,停用词表随便找一份常见的 1000 词版本即可。词表固定是整个实验里最先必须完成的一步,因为它决定三套模型的特征空间。下面这段代码同时输出词表和每个样本的 token id:

# text_utils.py import jieba from collections import Counter def load_lines(path): pairs = [] with open(path, encoding="utf-8") as f: for line in f: label, text = line.rstrip("\n").split("\t", 1) pairs.append((int(label), text)) return pairs def tokenize(pairs, stopwords): result = [] for label, text in pairs: words = [w for w in jieba.lcut(text) if w not in stopwords and w.strip()] result.append((label, words)) return result def build_vocab(tokenized, min_count=5, max_size=20000): counter = Counter() for _, words in tokenized: counter.update(words) vocab = {"<PAD>": 0, "<OOV>": 1} for w, c in counter.most_common(max_size - 2): if c < min_count: break vocab[w] = len(vocab) return vocab

逻辑说明:load_lines按第一个\t切成 (标签, 原文),tokenize做分词并过滤停用词和空串,build_vocab把词表压在前 20000 个词内。参数方面,min_count=5是经验值,太小会让低频噪音词进入后面的 PMI 计算,太大又会在短文本数据上制造大量 OOV;max_size=20000是为 TextGCN 的单位矩阵特征准备的,词表翻倍,单位矩阵和邻接矩阵的内存会成倍上涨,第 3 章会展开说这个约束。

注意:这里词表用的是 train+dev+test 全部文本。TextGCN 论文本身就把测试文档放进图里,因此这步对 TextGCN 是标准操作;但如果同时要对比 TextING 和 LEAM,建议再单独保存一份只用 train 构建的词表,避免答辩被问“测试集有没有泄漏进词表”时答不上来。

2.3 conda 创建 Python 环境与依赖锁定

复现这三套模型不需要很新的 CUDA,torch 2.x 的 CPU 版本也能跑完全部实验,只是慢一些。我用 conda 单独建环境,避免把本机的 Python 环境搅乱:

conda create -n nlp_hw python=3.10 -y conda activate nlp_hw pip install torch==2.1.2 --index-url https://download.pytorch.org/whl/cu118 pip install torch-geometric gensim scikit-learn jieba

如果用的是 VSCode,记得在命令面板里执行 “Python: Select Interpreter”,选中nlp_hw这个 conda env,否则终端里能 import 的包在编辑器里照样报 ModuleNotFoundError,这一步是 python 环境配置里卡住人最多的位置。依赖版本和用途如下:

包名版本建议用途
torch2.1.x三套模型的张量运算与自动求导
torch-geometric2.4+TextGCN 的 GCNConv、TextING 的 GatedGraphConv 和 Batch
gensim4.3+训练 Word2Vec 词向量,供 TextING 和 LEAM 使用
scikit-learn1.3+TF-IDF 计算与评估指标
jieba0.42+中文分词

torch-geometric 的安装有个传统坑:老版本依赖 torch-scatter、torch-sparse 等扩展包,直接 conda 装容易编译报错。torch-geometric 2.4 之后多数算子已经内置,先装 torch 再装 torch-geometric,顺序不要反。

2.4 用 gensim 在训练集上练一份 300 维词向量

TextGCN 不需要预训练词向量,它用单位矩阵当节点特征,语义靠图结构传递。TextING 和 LEAM 则必须有词向量做节点初始特征和标签表示。期末项目里最可控的做法不是下载外部词向量,而是就在训练集上自己训一份:

# train_w2v.py from gensim.models import Word2Vec from text_utils import load_lines, tokenize train_pairs = load_lines("data/train.txt") stopwords = set(open("stopwords.txt", encoding="utf-8").read().split()) train_tok = tokenize(train_pairs, stopwords) sentences = [w for _, w in train_tok] w2v = Word2Vec(sentences, vector_size=300, window=5, min_count=2, workers=8, epochs=10) w2v.save("w2v_300.model")

这里的参数含义:vector_size=300是主流默认值,太小装不下语义,太大对 LEAM 只会增加过拟合风险;min_count=2保留足够多的低频词来压低 OOV 比例;epochs=10对中小语料够用,再多容易把向量拟合到训练集的词汇分布上。有了词表和词向量,第 3 章到第 5 章的模型代码就有了统一输入。

3. TextGCN 复现:文档-词异构图、PMI 邻接与两层 GCN

3.1 为什么 TextGCN 是传导式的:整张图一次建好

TextGCN 的核心是构造一张包含“词节点 + 文档节点”的异构图。文档和词之间有边,权重是 TF-IDF;词和词之间有边,权重是正的点互信息 PMI。关键设计在于文档节点同时包含训练集、验证集和测试集,训练时只有训练节点带标签,却让所有节点在图里一起参与消息传递。这种“测试数据已经躺在图里”的学习方式就是传导式(transductive)。

传导式带来的工程后果很直接:拿到一篇训练时没见过的新文本,不能说“直接推理”,必须把新文档以及它包含的词节点挂到原图上,或者单独跑一次只含新文本的小图前向。期末答辩时把这一点讲清楚,比背模型结构更能证明你真的理解了 TextGCN,这也是它和 TextING 最本质的分界。

3.2 邻接矩阵构建:TF-IDF 边权重与正 PMI 词边代码

构建目标是产出两个集合:边起点、边终点、边权重。用 scipy 的 coo_matrix 组装稀疏矩阵,再转成 PyG 的 edge_index 和 edge_weight。TF-IDF 边直接让 sklearn 的 TfidfVectorizer 输出稀疏矩阵,最省事:

# build_graph.py import math import torch from collections import Counter from sklearn.feature_extraction.text import TfidfVectorizer def doc_word_edges(docs, vocab, offset): # docs: 已经分词并用空格拼接的文本列表,长度 = 文档数 vec = TfidfVectorizer(vocabulary=vocab, token_pattern=r"\S+", lowercase=False) tfidf = vec.fit_transform(docs) # (n_docs, n_vocab) coo = tfidf.tocoo() # 文档节点 id 从 0 开始;词节点 id 整体偏移 offset,避开文档节点 return coo.row, coo.col + offset, coo.data

PMI 词边稍微麻烦。PMI 的定义是log(p_ij / (p_i * p_j)),其中p_ij是两个词在滑动窗口内共现的概率,p_i是单词出现概率。论文只保留 PMI 大于 0 的边,因为负 PMI 在归一化后基本是噪音:

def pmi_edges(sentences, vocab, offset, window=20, min_count=5): word_freq, pair_freq, total_pairs = Counter(), Counter(), 0 for sent in sentences: sent = [w for w in sent if w in vocab] for i in range(len(sent)): word_freq[sent[i]] += 1 seen = set() # 同一窗口内同一对词只计一次,防止重复共现虚高 for j in range(i + 1, min(i + window, len(sent))): pair = (sent[i], sent[j]) if sent[i] <= sent[j] else (sent[j], sent[i]) if pair in seen: continue seen.add(pair) pair_freq[pair] += 1 total_pairs += 1 rows, cols, weights = [], [], [] for (a, b), count in pair_freq.items(): if min(word_freq[a], word_freq[b]) < min_count: continue p_ab = count / total_pairs p_a = word_freq[a] / sum(word_freq.values()) p_b = word_freq[b] / sum(word_freq.values()) pmi = math.log2(p_ab / (p_a * p_b)) if pmi > 0: rows += [vocab[a] + offset, vocab[b] + offset] # 无向图,补反向边 cols += [vocab[b] + offset, vocab[a] + offset] weights += [pmi, pmi] return rows, cols, weights

两个函数里的offset都是文档节点总数,因为节点顺序被固定为“文档在前、词在后”。三个参数按数据调:window=20是 TextGCN 论文原值,短文本可以降到 10,长文本可以放到 25;min_count=5过滤掉只出现几次的偶发词,这类词的 PMI 经常虚高;只保留正 PMI 是论文做法,负边参与归一化反而压低有效边的权重。

3.3 两层 GCN 前向与单位矩阵特征:没有预训练向量也能跑

TextGCN 的节点初始特征就是单位矩阵,每个节点一个 one-hot 向量。模型没有任何语义先验,语义完全靠图的边结构和标签传播两者共同注入,这也是它不需要预训练词向量的原因。实现直接用 PyG 的 GCNConv:

# textgcn.py import torch.nn.functional as F from torch_geometric.nn import GCNConv class TextGCN(torch.nn.Module): def __init__(self, num_nodes, hidden=200, num_classes=10): super().__init__() self.conv1 = GCNConv(num_nodes, hidden) self.conv2 = GCNConv(hidden, num_classes) def forward(self, x, edge_index, edge_weight): x = F.relu(self.conv1(x, edge_index, edge_weight)) return self.conv2(x, edge_index, edge_weight)

GCNConv内部默认加自环并按度做对称归一化,所以外部不用再手工处理行列归一化。edge_weight不传时全部为 1,这里必须传前一步算好的 TF-IDF 和 PMI 权重,否则边权信息全部丢失,准确率会显著下降。

训练前先构造整图的 Data 对象:

N = len(vocab) + len(docs) data = Data(x=torch.eye(N), edge_index=torch.tensor([rows, cols], dtype=torch.long), edge_weight=torch.tensor(weights, dtype=torch.float))

torch.eye(N)是稠密矩阵,N 在两万五左右时会占 2.5GB 内存。GPU 显存吃紧时把第 2 章的max_size降到 10000,或者直接先用 CPU 跑通再换 GPU。<PAD><OOV>这两个特殊节点没有连边,在图里是孤立节点,不会影响消息传递。

3.4 训练循环与 TextGCN 的四个关键超参数

训练时每个 epoch 只做一次全图前向,损失只算训练文档节点:

optimizer = torch.optim.SGD(model.parameters(), lr=0.02, weight_decay=1e-4) for epoch in range(200): model.train() out = model(data.x, data.edge_index, data.edge_weight) loss = F.cross_entropy(out[train_idx], train_y) optimizer.zero_grad() loss.backward() optimizer.step() if epoch % 20 == 0: model.eval() pred = out[dev_idx].argmax(1) print(epoch, round(loss.item(), 4), round((pred == dev_y).float().mean().item(), 4))

out[train_idx]能这样索引的前提是节点顺序固定:前若干个是训练文档,接着是 dev、test、词节点。顺序在构图时必须记死,否则标签对不上节点。lr=0.02用 SGD 是论文默认值,换成 Adam 后学习率要降到 1e-3 量级,否则前期 loss 抖动剧烈。

参数推荐值影响
GCN 层数2超过 2 层容易过平滑,节点表示趋同
hidden200小于 100 时图语义承载不够
PMI window20越大共现边越多,越小图越稀疏
dropout0.5加在两层 GCN 输出之间
weight_decay1e-4图模型对过拟合比序列模型更敏感

两个高频坑:一是 TF-IDF 边和 PMI 边拼接时顺序弄混,edge_weight 是按边顺序对应的,拼之前记下每部分边数;二是换了 session 后忘记重载 edge_index,导致图结构与词表对不上,训练 loss 正常但准确率停在类别先验水平。

4. TextING 复现:每篇文档独立词图上的归纳式消息传递

4.1 TextING 的建模假设和 TextGCN 差在哪

TextGCN 是全语料一张图,TextING 反过来:一篇文档一张图。文档里的每个词是一个节点,用滑动窗口的共现关系连边,然后在这个小图上做若干轮 Gated 消息传递,最后对全部节点特征做池化得到文档表示。因为训练和测试的图彼此独立,TextING 是归纳式(inductive)模型,新样本进来直接建自己的小图就能出结果,这是它和 TextGCN 最本质的区别。

这个差异直接决定工程取舍。TextGCN 每来一批新数据都要面对“挂回旧图还是整图重建”的选择,TextING 完全没有这个问题,天然适合线上单条预测。代价是 TextING 看不到跨文档的词共现信息,在短文本数据集上通常比 TextGCN 低 1~2 个百分点,换来的是可部署性。期末作业里把这段取舍写进报告,比“我们用了图神经网络”有说服力得多。

4.2 单文档词图构建与 Gated 消息传递代码

先用 PyG 的 Data 把一篇文档的 token id 序列转成图:

# texting_graph.py import torch from torch_geometric.data import Data def doc_to_graph(token_ids, window=3): n = len(token_ids) src, dst = [], [] for i in range(n): for j in range(i + 1, min(i + window, n)): src += [i, j] dst += [j, i] edge_index = torch.tensor([src, dst], dtype=torch.long) return Data(x=torch.tensor(token_ids, dtype=torch.long), edge_index=edge_index)

window=3表示每个词只和它后面的两个词连边,并补反向边。窗口越小图越稀疏,消息传递越局部;窗口超过 5 后,短文本的图基本变成全连接,池化出来的表示和均值向量没区别,图结构的信息就浪费了。

模型主体直接用 PyG 的 GatedGraphConv,它封装了 Gated GNN 的消息传递循环:

# texting.py import torch import torch.nn as nn import torch.nn.functional as F from torch_geometric.nn import GatedGraphConv class TextING(nn.Module): def __init__(self, embed, hidden=300, steps=2, num_classes=10, dropout=0.5): super().__init__() self.embed = nn.Embedding.from_pretrained(embed, freeze=False, padding_idx=0) self.proj = nn.Linear(embed.size(1), hidden) self.ggnn = GatedGraphConv(hidden, steps) self.dropout = nn.Dropout(dropout) self.fc = nn.Linear(hidden, num_classes) def forward(self, batch, n_graphs, pad_mask): x = self.dropout(self.proj(self.embed(batch.x))) # (总节点数, hidden) x = self.ggnn(x, batch.edge_index) x = x.clone() x[~pad_mask] = -1e9 # padding 词不参与池化 out = torch.full((n_graphs, x.size(1)), -1e9, device=x.device) out.scatter_reduce_(0, batch.batch.unsqueeze(1).expand_as(x), x, reduce="amax", include_self=False) return self.fc(self.dropout(out))

GatedGraphConv(hidden, steps)steps是消息传递轮数,短文本 2 轮就够,3 轮以上信息已经在全图来回多趟,边际收益很小。batch.batch是 PyG 的 Batch 在拼接多个图时自动生成的“节点属于哪篇文档”的索引,池化就是按它分桶取最大值。pad_mask在数据加载时构造,形状是总节点数,padding 位置为 False。把 padding 词替换成-1e9而不是乘 0,是因为乘 0 会让 max 池化在整行特征都为负时误选 padding 位置的 0。

注意:漏掉 pad_mask 是 TextING 复现里最常见的错误。padded 词没有连边,经过消息传递后仍带着投影层的输出,scatter max 会把它们当作正常节点,典型表现是训练 loss 正常下降、验证集准确率长期不动。

4.3 Batch 拼接、词向量初始化和三个必调参数

训练时把所有文档一次性 Batch 成一张大图,比循环单篇快一个数量级:

from torch_geometric.data import Batch def collate_graphs(graph_list): return Batch.from_data_list(graph_list)

实际 data loader 会把每批文档先 padding 成等长,再逐一转图,最后Batch.from_data_list拼起来。文档长度不齐没关系,Batch 会自动记录每个子图的边界。词向量用第 2 章训好的 Word2Vec 构建 embedding 矩阵,freeze=False表示训练时微调词向量;数据量在一两万条时微调收益明显,数据量再大为了防止过拟合改成freeze=True

参数推荐值说明
steps2消息传递轮数,短文本 2 轮足够
window3共现窗口,过大会让图接近全连接
hidden300词向量维度和 hidden 不一致时由 proj 对齐
dropout0.5加在投影层和输出层

优化器和 TextGCN 完全不同:这里用 Adam,lr 取 1e-3,epoch 数 30~50 就够收敛,不需要 200 轮。TextING 没有单位矩阵那样的大稠密张量,同样数据量下内存占用明显低于 TextGCN,CPU 也能跑得舒服。

5. LEAM 复现:标签嵌入引导的文本 attention,不建图也能竞速

5.1 LEAM 的核心思想:让标签语义参与词的加权

TextGCN 和 TextING 都是图模型,LEAM 走的是完全不同的路线。它认为一个词对分类结果的贡献,取决于它和各个类别标签的语义接近程度。做法是把每个类别的标签文本映射成词向量,和输入文本的词向量互算相似度,相似度高的词拿到更大的 attention 权重,最后把加权后的文本表示和标签向量做内积得到分类分数。

这套设计的好处是模型轻、训练快、没有构图逻辑。对情感分析、新闻分类这类标签语义明确的场景,LEAM 经常能追平甚至超过图模型,而且它只需要词向量,和 TextING 共用第 2 章训练好的那份即可。期末大作业里把 LEAM 当作“无图基线”和两个图模型对照,可以直接回答“图结构到底带来了多少增益”这个问题。

5.2 LEAM 前向传播:CNN 特征、词-标签相似度与注意力加权代码

先用第 2 章的词表和词向量,把每个类别的标签文本转成向量均值,作为标签嵌入的初始化:

# leam.py import torch import torch.nn as nn import torch.nn.functional as F def label_embed_init(word_segs_by_class, w2v, dim=300): # word_segs_by_class[c] 是第 c 类全部训练文本分词后的词列表,已拼接 init = [] for segs in word_segs_by_class: vecs = [torch.tensor(w2v.wv[w]) for w in set(segs) if w in w2v.wv] init.append(torch.stack(vecs).mean(0)) return torch.stack(init) # (C, dim)

逻辑说明:每个类的标签嵌入是该类训练文本全部词向量的平均。比随机初始化收敛更快,也避免了“标签文本本身太短,平均出来信息量不够”的问题。dim=300必须和词向量维度一致。

模型前向走四步:查词向量、过一维卷积提取 n-gram 特征、和标签嵌入算相似度得到 attention、加权求和后与标签做内积:

class LEAM(nn.Module): def __init__(self, embed, label_embed, hidden=300, conv_kernel=3, temperature=2.0, dropout=0.5): super().__init__() self.embed = nn.Embedding.from_pretrained(embed, freeze=False, padding_idx=0) self.conv = nn.Conv1d(hidden, hidden, conv_kernel, padding=conv_kernel // 2) self.dropout = nn.Dropout(dropout) self.label_embed = nn.Parameter(label_embed) # (C, d) self.temp = temperature def forward(self, x): # x: (B, L) 词 id w = self.dropout(self.embed(x)) # (B, L, d) w = F.relu(self.conv(w.transpose(1, 2))) # (B, d, L) w = w.transpose(1, 2) # (B, L, d) s = torch.matmul(w, self.label_embed.t()) / self.temp # (B, L, C) alpha = F.softmax(s.max(dim=-1).values, dim=1).unsqueeze(-1) doc = (w * alpha).sum(dim=1) # (B, d) return torch.matmul(doc, self.label_embed.t())

s.max(dim=-1)取每个词对所有标签相似度的最大值,代表“这个词和哪个标签最像”;softmax 在词维上归一化得到 attention。temperature=2.0让 attention 分布更平滑,避免少数高频词一枝独秀;调小会让注意力更“硬”,风险更大。用nn.Parameter包装标签嵌入使标签向量随训练更新,这是 LEAM 区别于普通 attention 机制的关键。

注意:label_embed.t()参与 matmul,要求 label_embed 是 float 且和词向量同维度。最常见的报错是把标签嵌入初始化成了整数张量,或者类别数写成了 batch size,导致维度对不上。

5.3 LEAM 的三个调参点和最容易过的拟合关卡

参数推荐值说明
conv_kernel3卷积窗口,可以并列多个 kernel 模拟 n-gram
temperature2.0attention 平滑度,越大越分散
dropout0.5对过拟合最敏感的参数,没有之一
lr (Adam)1e-3相比两个图模型更容易欠拟合,lr 可以略大

LEAM 的模型体量最小,最容易出现的是过拟合而不是欠拟合:词表 2 万、训练数据才 1 万条时,embedding 层占掉绝大部分参数。应对手段是freeze=True固定词向量,只更新卷积和标签向量;或者把 hidden 从 300 压到 128。判断是不是过拟合,看训练准确率和 dev 准确率的差,超过 8 个百分点就该收紧模型容量。

到这里,三套模型的输入输出已经统一成(batch, num_classes)的 logits,第 6 章可以直接在一个脚本里对齐评估。

6. 三种文本分类方法的统一评估与 logits 平均值集成

6.1 用同一份测试脚本对齐三套模型的输出

三套模型训练完成后,都保存验证集最优的 checkpoint,在测试集上输出 logits。只要类别顺序来自同一份 train.txt,三份 logits 就可以直接对比:

# evaluate.py import numpy as np import torch from sklearn.metrics import accuracy_score, f1_score, classification_report logits = { "textgcn": np.load("out_textgcn.npy"), "texting": np.load("out_texting.npy"), "leam": np.load("out_leam.npy"), } for name, z in logits.items(): pred = z.argmax(1) print(name, "acc=%.4f" % accuracy_score(y_test, pred), "macro-F1=%.4f" % f1_score(y_test, pred, average="macro")) print(classification_report(y_test, logits["textgcn"].argmax(1)))

argmax(1)在第二维取最大值的下标。多分类评估里 macro-F1 比 accuracy 更能反映少数类表现,TextGCN 在类别不平衡的数据上经常 accuracy 高、macro-F1 低,因为它会把少数类节点的表示往多数类方向拉。

6.2 复现过程中最常见的六类问题定位

现象原因处理方式
TextGCN CUDA 显存不足torch.eye 稠密矩阵太大词表上限降到 10000,或改用 CPU 训练
TextGCN 准确率像随机猜edge_weight 没传或索引错位检查边数是否和 weight 长度一致
TextING 验证集不涨padded 词参与最大值池化池化前把 padding 位置替换成 -1e9
TextING 收敛慢window 太大、图太密window 降到 3,steps 降到 2
LEAM 训练集好测试集差embedding 层参数过多freeze 词向量,hidden 压到 128
三模型结果无法对齐词表或类序不一致统一用 train.txt 生成的同一份 vocab

6.3 一个立竿见影的集成技巧:三个模型 logits 平均

图模型和注意力模型的错误分布通常不重叠,把它们的结果做概率平均,大概率能拿到比最好单模型更高的准确率。注意先做 softmax 再平均,不要直接平均原始 logits,因为 TextGCN 的 logits 方差比其他两个模型大不少,直接平均会被它的量级带偏:

probs = [torch.softmax(torch.tensor(z), dim=-1) for z in logits.values()] avg = torch.stack(probs).mean(0) print("ensemble acc=%.4f" % accuracy_score(y_test, avg.argmax(1).numpy()))

如果想让集成再进一步,可以给三个模型配权重,用验证集做小规模搜索。一个常见的经验起点是 TextGCN 0.4、TextING 0.3、LEAM 0.3,数据偏短文本时 LEAM 的权重可以再拉高。保存 logits 时务必用同一个测试集顺序和同一个类标签顺序,这里最隐蔽的错误是两份 npy 文件的行顺序不一致,会让集成结果反而低于最好的单模型。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/11 19:09:16

共享单车租赁预测:从数据清洗到随机森林与SVM建模全流程

简介&#xff1a;共享单车租赁数量预测是交通出行类数据挖掘的典型场景&#xff0c;项目使用Python完成数据分析和可视化&#xff0c;并以随机森林、支持向量机两种模型进行租赁量回归预测。代码约300行&#xff0c;既涵盖Pandas、NumPy的数据清洗与特征预处理&#xff0c;也包…

作者头像 李华
网站建设 2026/9/11 19:06:15

插件、MCP、Skill 的区别?

&#x1f4cc; 一句话总结 《插件、MCP、Skill 的区别&#xff1a;工具封装&#xff08;平台生态&#xff09;、开放协议&#xff08;AI 的 USB-C&#xff09;、知识指令&#xff08;Markdown 编码的领域智慧&#xff09;——大模型能力扩展的三层演进》 详细总结 一、从一个…

作者头像 李华
网站建设 2026/9/11 19:06:12

喷涂机器人与人工喷涂对比:制造业喷涂自动化转型决策参考

一、制造业喷涂作业的转型背景近年来&#xff0c;制造业面临前所未有的挑战&#xff1a;用工成本持续攀升、熟练喷涂工人短缺、产品质量要求不断提高、环保监管日趋严格。在这样的背景下&#xff0c;越来越多的企业开始思考一个核心问题&#xff1a;喷涂机器人 vs 人工喷涂&…

作者头像 李华
网站建设 2026/9/11 19:02:32

基于AT89C51的烟雾浓度与温度检测火灾报警系统设计

简介&#xff1a;这是一份基于AT89C51单片机的温度烟雾火灾报警系统设计方案&#xff0c;面向单片机初学者与电子设计竞赛参赛者。系统具备火情探测、灯光报警、蜂鸣器报警和阈值调节等核心功能。烟雾采集选用MQ-2传感器&#xff0c;输出模拟信号&#xff0c;经ADC0832转换后送…

作者头像 李华
网站建设 2026/9/11 19:02:23

JAVA毕设选题推荐:基于 SpringBoot 的烘焙电商平台的设计与实现 基于 SpringBoot 技术的甜点商城系统【附源码、mysql、文档、调试+代码讲解+全bao等】

博主介绍&#xff1a;✌️码农一枚 &#xff0c;专注于大学生项目实战开发、讲解和毕业&#x1f6a2;文撰写修改等。全栈领域优质创作者&#xff0c;博客之星、掘金/华为云/阿里云/InfoQ等平台优质作者、专注于Java、小程序技术领域和毕业项目实战 ✌️技术范围&#xff1a;&am…

作者头像 李华