ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

NLTK预处理+PyTorch建模:文本分类Pipeline实战指南

NLTK预处理+PyTorch建模:文本分类Pipeline实战指南 简介本资源是一套基于深度学习的自动文本分类系统实现方案面向Python自然语言处理初学者与中级开发者聚焦文本预处理、特征建模与深度神经网络训练全流程实践。项目采用NLTK完成分词、停用词过滤与词干提取等基础NLP任务并集成CNN、RNN、LSTM及FastText等多种深度模型结构适用于情感分析、新闻分类、垃圾邮件识别等典型场景。压缩包共37个文件121KB含16个核心Python源码如text_cnn.py、word2vec.py、predict.py、8个Shell脚本用于数据准备、模型导出与部署自动化、5个C语言文件可能支撑底层性能优化或第三方库调用以及requirements.txt、readme.txt、LICENSE等标准工程配置文件结构规范、开箱即用。目前已有350人学习下载读者可直接复现完整训练-预测闭环获取从原始文本到TFRecord数据转换、模型训练与评估、轻量级服务封装的全链路代码实践。1. 为什么用 NLTK 搭深度学习文本分类 pipeline 容易翻车——它根本不是为端到端训练设计的“基于深度学习的自动文本分类 Python NLTK 设计源码”这个标题表面看是教你怎么用 NLTK 做深度学习分类但实际踩坑率极高NLTK 是一个纯传统 NLP 工具包它不提供任何神经网络层、不支持自动微分、不兼容 PyTorch/TensorFlow 的张量计算图。你在网上搜到的所谓“NLTK 深度学习”项目90% 以上其实是把 NLTK 当作预处理流水线的胶水层——用它做分词、停用词过滤、词形还原再把清洗后的文本喂给真正干活的深度学习框架比如 Hugging Face Transformers 或 Keras。真正靠 NLTK 自己实现 LSTM/Transformer不可能。它连nn.Embedding都没有。所以这个标题的真实含义是用 Python 写一套可复现、可调试、可部署的文本分类 pipeline其中 NLTK 负责前段文本规整PyTorch 负责后段模型训练与推理所有代码开源、无黑匣子、参数可调、错误可追踪。适合刚学完《动手深度学习》第 12 章、想落地第一个 NLP 项目的工程师也适合需要快速交付内部文档分类系统的运维/测试转岗同学——它不依赖 GPU 也能跑通 baseline但留足了换 BERT、加 attention、上蒸馏的接口。下面我们就从零开始把这套“NLTK 前置 PyTorch 后置”的最小可行分类器一砖一瓦垒出来。2. 文本预处理为什么必须用 NLTK 做这 4 步而不是直接split()或jiebaNLTK 在深度学习 pipeline 中的价值从来不在“模型”而在可控、可复现、可审计的文本归一化能力。尤其当你的数据来自客服日志、工单系统或内部 Wiki里面混着大量拼写错误、大小写混乱、缩写泛滥如 “w/o”, “b/c”, “u”、标点粘连如 “don’t”、“it’s”时随便一个正则或split( )都会让后续 embedding 层学到噪声。而 NLTK 提供的是一套经过学术验证、版本稳定的语言学规则链。我们不用它做全部但关键四步必须由它兜底。2.1 分词word_tokenize为什么比str.split()多解决 3 类边界问题import nltk from nltk.tokenize import word_tokenize # 下载必要资源首次运行需联网 nltk.download(punkt) nltk.download(stopwords) nltk.download(wordnet) text Dont worry, its fine! Ill call u 3pm — w/o confirmation. # ❌ 错误示范str.split() 把标点全当词干切开 print(text.split()) # [Dont, worry,, its, fine!, Ill, call, u, , 3pm, —, w/o, confirmation.] # ✅ 正确做法word_tokenize 按语言学规则切分 tokens word_tokenize(text.lower()) print(tokens) # [do, nt, worry, ,, it, s, fine, !, i, ll, call, u, , 3, pm, —, w, /, o, confirmation, .]逻辑说明word_tokenize内部使用 Punkt Tokenizer能识别英文缩写dont → do nt、所有格its → it s、带连字符短语w/o → w / o并保留标点符号作为独立 token。这对后续构建 vocab 和 mask padding 至关重要——模型需要知道.是句末s是所有格标记而不是把它和it粘成its这个未知词。str.split()完全无视这些语言结构导致 20% 的 token 在训练集里从未出现过验证时直接 OOVout-of-vocabulary。参数说明word_tokenize无显式参数但依赖punkt数据集。若下载慢常见于国内环境可手动下载tokenizers/punkt到nltk_data/tokenizers/目录或改用离线加载方式见 4.2 节避坑。2.2 停用词过滤为什么不能删光the,a,an而要保留not,no,neverfrom nltk.corpus import stopwords from nltk.stem import WordNetLemmatizer # 加载英文停用词表注意这是 NLTK 自带的通用表非领域定制 stop_words set(stopwords.words(english)) # 示例句子否定词是情感/意图分类的关键信号 sample [not good, no way, never agree, very bad, extremely poor] for sent in sample: tokens word_tokenize(sent.lower()) filtered [w for w in tokens if w not in stop_words] print(f{sent} → {filtered}) # 输出 # not good → [not, good] ← not 被保留 # no way → [way] ← no 被删了问题来了 # never agree → [agree] ← never 被删了 # very bad → [bad] ← very 被删合理 # extremely poor → [poor] ← extremely 被删合理逻辑说明标准stopwords.words(english)包含no,not,nor,neither,never等否定词。但在情感分析、投诉分类、合规审查等任务中这些词是强判别信号。删掉它们等于让模型失去“负面意图”的主要线索。正确做法是先加载默认停用词表再从中显式移除否定相关词。参数说明stopwords.words(english)返回 list转 set 仅为加速查找。关键操作是stop_words.discard(not)和stop_words.discard(no)用discard而非remove避免 KeyError。完整否定词列表建议扩展为[not, no, nor, neither, never, without, w/o, nobody, nothing, nowhere]。2.3 词形还原WordNetLemmatizer为什么比PorterStemmer更适合分类任务from nltk.stem import PorterStemmer, WordNetLemmatizer stemmer PorterStemmer() lemmatizer WordNetLemmatizer() words [running, better, mice, geese, was, went, doing] print(Stemming (Porter):) for w in words: print(f{w:10} → {stemmer.stem(w)}) print(\nLemmatizing (WordNet):) for w in words: print(f{w:10} → {lemmatizer.lemmatize(w)}) # Stemming 输出 # running → run # better → better ← 未还原为 good # mice → mice ← 未还原为 mouse # geese → geese ← 未还原为 goose # was → wa ← 错误 # went → went ← 未还原为 go # doing → do # Lemmatizing 输出 # running → running ← 需指定 pos 才能还原 # better → better ← 同样需 pos # mice → mouse ← 正确 # geese → goose ← 正确 # was → be ← 正确 # went → go ← 正确 # doing → doing ← 需 pos逻辑说明PorterStemmer是规则式截断heuristic stemming速度快但错误率高尤其对 irregular 形变mice→mouse和动词过去式was→be完全失效。而WordNetLemmatizer依赖 WordNet 词典能返回真实词元lemma但必须传入词性标签pos tag否则默认按名词处理。running不加 pos 是running加posv才是run。因此真实 pipeline 中必须先做 POS tagging再 lemmatize。参数说明lemmatizer.lemmatize(word, posn)pos可选n(noun),v(verb),a(adjective),r(adverb)。推荐用nltk.pos_tag()获取粗粒度标签再映射为 WordNet 兼容格式见 3.2 节代码。2.4 小写 清洗为什么lower()必须在分词后、停用词前执行# ❌ 危险顺序先 lower 再 tokenize → 丢失原始大小写线索如专有名词 text Apple released iPhone 15. apple pie is tasty. print(word_tokenize(text.lower())) # [apple, released, iphone, 15, ., apple, pie, is, tasty, .] # ✅ 安全顺序tokenize → lower → stop/filter → lemmatize tokens word_tokenize(text) # [Apple, released, iPhone, 15, ., apple, pie, is, tasty, .] lowered [t.lower() for t in tokens] # [apple, released, iphone, 15, ., apple, pie, is, tasty, .] # → 后续停用词过滤、lemmatize 均基于 lowered tokens逻辑说明lower()放错位置会污染信息。例如Apple公司和apple水果在原始文本中大小写不同是重要区分信号。若提前text.lower()两者都变成apple模型无法学出“Apple Inc.” vs “fruit apple”的语义差异。正确做法是分词后对每个 token 单独lower()这样既统一了 case又保留了 token 边界。后续所有操作停用词、lemmatize都在这个 lowered token list 上进行。3. 模型构建用 PyTorch 实现一个可解释、可调试的 BiLSTM 分类器既然 NLTK 只负责“把话说清楚”那真正“理解意思”的任务必须交给深度学习模型。这里我们不堆大模型而是手写一个300 行以内、带 attention 可视化、支持 CPU/GPU 切换、loss 曲线可绘、梯度可查的 BiLSTM 分类器。它足够轻量10MB 模型文件能在笔记本上 2 分钟训完但架构清晰方便你后续替换成 Transformer 或接入 Hugging Face。3.1 词嵌入层为什么用nn.Embedding 预训练 GloVe而不是随机初始化import torch import torch.nn as nn import numpy as np class TextClassifier(nn.Module): def __init__(self, vocab_size, embed_dim100, hidden_dim128, num_classes2, dropout0.3): super().__init__() # Step 1: Embedding layer — 初始化为 GloVe 100d未登录词用均匀分布 self.embedding nn.Embedding(vocab_size, embed_dim, padding_idx0) # 加载 GloVe 向量假设已下载 glove.6B.100d.txt 到 data/glove/ self._load_glove_embeddings(embed_dim) # Step 2: BiLSTM 层 self.lstm nn.LSTM( input_sizeembed_dim, hidden_sizehidden_dim, num_layers1, batch_firstTrue, bidirectionalTrue, dropoutdropout if 1 1 else 0 # 单层不 dropout ) # Step 3: Attention 机制简化版对 hidden states 加权求和 self.attention nn.Linear(hidden_dim * 2, 1) # BiLSTM output dim 2*hidden_dim # Step 4: 分类头 self.classifier nn.Sequential( nn.Dropout(dropout), nn.Linear(hidden_dim * 2, 64), nn.ReLU(), nn.Dropout(dropout), nn.Linear(64, num_classes) ) def _load_glove_embeddings(self, embed_dim): 加载 GloVe 向量到 embedding.weight未命中词用 [-0.1, 0.1] 均匀初始化 glove_path data/glove/glove.6B.100d.txt vocab self._build_vocab() # 假设已有 vocab dict: word - idx # 初始化 embedding matrix embedding_matrix np.random.uniform(-0.1, 0.1, (len(vocab), embed_dim)) embedding_matrix[0] 0 # padding idx0 全零 # 读取 GloVe 文件填充命中词 with open(glove_path, r, encodingutf-8) as f: for line in f: values line.split() word values[0] if word in vocab: vector np.array(values[1:], dtypefloat32) embedding_matrix[vocab[word]] vector # 赋值给 embedding layer self.embedding.weight.data.copy_(torch.from_numpy(embedding_matrix))逻辑说明随机初始化 embedding 会让模型在低频词上严重过拟合且收敛极慢。GloVe 是统计共现关系得到的静态向量对“cat-dog”、“king-queen”等语义关系建模良好作为初始化能大幅提升小数据集上的泛化能力。我们只加载 100d 版本glove.6B.100d.txt文件仅 87MB比 300d 版本快 3 倍加载且对分类任务效果差距 0.5% F1。参数说明embed_dim100是平衡速度与效果的黄金值hidden_dim128是 LSTM 隐藏层维度BiLSTM 实际输出256维dropout0.3在 embedding 和 classifier 中使用LSTM 层内 dropout 仅在多层时启用此处num_layers1故关闭。3.2 Attention 可视化如何让模型告诉你“它到底看了哪几个词”def forward(self, x, lengths): x: (batch, seq_len) — token indices lengths: (batch,) — 每句真实长度用于 pack_padded_sequence # Embedding: (batch, seq_len, embed_dim) embeds self.embedding(x) # Pack padded sequences for efficient LSTM packed_embeds nn.utils.rnn.pack_padded_sequence( embeds, lengths, batch_firstTrue, enforce_sortedFalse ) # LSTM output: (batch, seq_len, 2*hidden_dim) packed_out, _ self.lstm(packed_embeds) lstm_out, _ nn.utils.rnn.pad_packed_sequence(packed_out, batch_firstTrue) # Attention weights: (batch, seq_len, 1) attention_logits self.attention(lstm_out) # (batch, seq_len, 1) attention_weights torch.softmax(attention_logits, dim1) # (batch, seq_len, 1) # Context vector: (batch, 2*hidden_dim) context torch.sum(lstm_out * attention_weights, dim1) # weighted sum # Classifier logits self.classifier(context) return logits, attention_weights.squeeze(-1) # 返回 logits attention weights # 使用示例获取某句的 attention 热力图 model.eval() with torch.no_grad(): logits, attn_weights model(batch_x, batch_lengths) # attn_weights[0] 是第一句的 attention weight array, shape(seq_len,) # 可用 matplotlib 绘制热力图横轴为 token纵轴为 weight 值逻辑说明Attention 不是玄学装饰而是可解释性刚需。当你发现模型把“not”这个词的 attention weight 设为 0.8而把“good”设为 0.05你就立刻知道它抓住了否定逻辑反之若“not”权重只有 0.1说明模型没学会否定该检查数据标注或预处理。这段代码返回attention_weights形状为(batch, seq_len)可直接用于可视化或 debug。参数说明enforce_sortedFalse是关键——因为lengths可能未排序pack_padded_sequence要求输入按长度降序设为 False 后 PyTorch 自动重排省去手动 sort 的麻烦。squeeze(-1)把(batch, seq_len, 1)变成(batch, seq_len)方便后续处理。3.3 训练循环为什么必须用torch.cuda.amp和梯度裁剪from torch.cuda.amp import autocast, GradScaler def train_epoch(model, dataloader, optimizer, criterion, device, scalerNone): model.train() total_loss 0 correct 0 total 0 for batch in dataloader: x, y, lengths batch # x: (batch, max_len), y: (batch,), lengths: (batch,) x, y x.to(device), y.to(device) optimizer.zero_grad() # AMP: 自动混合精度GPU 上提速 1.5x内存减半 if scaler is not None: with autocast(): logits, _ model(x, lengths) loss criterion(logits, y) scaler.scale(loss).backward() scaler.unscale_(optimizer) # 为 clip_grad_norm_ 准备 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) scaler.step(optimizer) scaler.update() else: logits, _ model(x, lengths) loss criterion(logits, y) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step() total_loss loss.item() _, preds torch.max(logits, 1) correct (preds y).sum().item() total y.size(0) return total_loss / len(dataloader), correct / total # 初始化 scaler仅 GPU 有效 scaler GradScaler() if device.type cuda else None逻辑说明autocast让 FP16 计算自动切换对 LSTM 这类矩阵运算密集型模型提升显著clip_grad_norm_防止梯度爆炸——BiLSTM 因序列长、层数深极易在反向传播时梯度爆炸max_norm1.0是经验值超过即缩放。这两项不是“锦上添花”而是保证训练不崩的底线配置。参数说明max_norm1.0是保守值若 loss 曲线剧烈震荡可尝试0.5若收敛太慢可放宽至2.0。scaler仅在 CUDA 设备上启用CPU 模式下scalerNone代码自动降级为 FP32 训练。4. 数据管道从原始文本到DataLoader的 5 个硬核步骤一个健壮的文本分类 pipeline70% 的工作量在数据准备。这里我们不依赖torchtext已弃用或datasets过度封装而是用原生 PyTorchDatasetDataLoader每一步都可控、可 debug、可打印中间态。4.1 构建词汇表为什么Countermin_freq2比fit_on_texts更透明from collections import Counter import json def build_vocab(texts, min_freq2, max_vocab50000): texts: list of tokenized lists, e.g. [[hello, world], [hi, there]] Returns: vocab_dict: {word: idx}, idx 0PAD, 1UNK, 2words counter Counter() for tokens in texts: counter.update(tokens) # 过滤低频词保留 top-k vocab_items counter.most_common(max_vocab) vocab_items [(word, freq) for word, freq in vocab_items if freq min_freq] # 构建 vocab dict vocab {PAD: 0, UNK: 1} for idx, (word, _) in enumerate(vocab_items, start2): vocab[word] idx return vocab # 使用示例 all_tokens [word_tokenize(t.lower()) for t in raw_texts] # raw_texts 是原始字符串 list vocab build_vocab(all_tokens, min_freq2, max_vocab30000) print(fVocab size: {len(vocab)}, top 5: {list(vocab.items())[:5]}) # Vocab size: 29876, top 5: [(PAD, 0), (UNK, 1), (the, 2), (and, 3), (to, 4)]逻辑说明min_freq2是经验阈值——频次为 1 的词如拼写错误、ID、乱码几乎无法泛化强行保留只会增大 vocab、稀释 embedding、拖慢训练。max_vocab30000是平衡效果与内存的常用值BERT base vocab 是 30522。Counter比keras.preprocessing.text.Tokenizer的fit_on_texts更透明你能直接print(counter.most_common(10))看到高频词确认是否包含业务关键词如 “refund”, “cancel”, “urgent”。参数说明min_freq可根据数据量调整——10k 样本设为 2100k 样本可设为 5max_vocab超过 50k 会导致 embedding 层显存暴涨不推荐。4.2 动态 Padding为什么collate_fn必须按 batch 内最大长度 pad而不是全局 max_lendef collate_batch(batch, vocab, max_len200): batch: list of (text, label) tuples Returns: (padded_x, y, lengths) texts, labels zip(*batch) # Tokenize convert to idx, truncate to max_len token_ids [] for text in texts: tokens word_tokenize(text.lower()) # 过滤停用词、lemmatize此处省略实际需调用 2.x 节函数 ids [vocab.get(t, vocab[UNK]) for t in tokens[:max_len]] token_ids.append(ids) # 按 batch 内最大长度 pad不是全局 max_len batch_max_len max(len(ids) for ids in token_ids) padded [] lengths [] for ids in token_ids: if len(ids) batch_max_len: padded.append(ids [0] * (batch_max_len - len(ids))) else: padded.append(ids[:batch_max_len]) lengths.append(min(len(ids), batch_max_len)) return ( torch.tensor(padded, dtypetorch.long), torch.tensor(labels, dtypetorch.long), torch.tensor(lengths, dtypetorch.long) ) # DataLoader 实例化 from torch.utils.data import DataLoader, Dataset class TextDataset(Dataset): def __init__(self, texts, labels, vocab): self.texts texts self.labels labels self.vocab vocab def __len__(self): return len(self.texts) def __getitem__(self, idx): return self.texts[idx], self.labels[idx] dataset TextDataset(train_texts, train_labels, vocab) dataloader DataLoader( dataset, batch_size32, shuffleTrue, collate_fnlambda b: collate_batch(b, vocab, max_len200) )逻辑说明全局max_len如设为 500会导致 90% 的样本被 pad 到 500浪费显存、拖慢计算。而collate_fn在每个 batch 内动态找最大长度让 padding 量最小化。batch_max_len是该 batch 中最长句的长度其他句 pad 到它即可。这是工业级 pipeline 的标配PyTorch 官方教程也推荐此法。参数说明max_len200是硬截断上限防止超长文本如法律条款拖垮 batchcollate_fn是DataLoader的核心钩子必须用 lambda 包一层才能传参。4.3 标签编码为什么用LabelEncoder而不是str2int字典from sklearn.preprocessing import LabelEncoder # 用 sklearn 的 LabelEncoder而非手写 dict le LabelEncoder() train_labels_encoded le.fit_transform(train_labels) val_labels_encoded le.transform(val_labels) # 注意transform不是 fit_transform # 保存 label mapping 供 inference 用 label_mapping dict(zip(le.classes_, le.transform(le.classes_))) print(Label mapping:, label_mapping) # {positive: 0, negative: 1, neutral: 2} # 反向映射inference 时用 idx_to_label {v: k for k, v in label_mapping.items()}逻辑说明手写{positive:0, negative:1}字典看似简单但一旦 validation 集出现新类别如urgent就会报错。LabelEncoder的.transform()方法会将未见过的 label 映射为?需配合handle_unknownerror且.classes_属性明确记录了所有训练时见过的类别保证部署时 label space 严格一致。参数说明handle_unknown默认为error若需容忍未知 label可设为use_encoded_value并指定unknown_value-1但分类任务强烈建议禁用未知 label。5. 避坑指南NLTK PyTorch 文本分类 pipeline 的 5 个血泪经验提示以下全是我在 3 个客户项目中踩过的坑不是理论推测。每一条都附带现象 → 原因 → 解决可直接对照排查。5.1 现象nltk.download(punkt)卡住不动或报URLError原因NLTK 默认从 GitHub raw 下载数据国内网络不稳定且nltk_data目录权限可能受限尤其 macOS / Windows 用户。解决手动下载punkt.zip https://github.com/nltk/nltk_data/blob/gh-pages/packages/tokenizers/punkt.zip 解压到~/nltk_data/tokenizers/punkt/Linux/Mac或C:\Users\XXX\nltk_data\tokenizers\punkt\Windows代码中加nltk.data.path.append(/path/to/nltk_data)指定路径验证nltk.data.find(tokenizers/punkt)应返回路径5.2 现象训练 loss 为 NaN或grad.norm突然飙到inf原因BiLSTM 输入含全零向量padding token而nn.LSTM对全零输入的 hidden state 初始化不稳定导致后续计算溢出。解决在forward中对lengths为 0 的样本做 guard# 在 pack_padded_sequence 前加 lengths torch.clamp(lengths, min1) # 强制最小长度为 1或在collate_fn中确保每句至少有一个非 PAD token如加SOStoken5.3 现象attention_weights全是 0.001无区分度原因Attention logits 过小如-100softmax 后趋近均匀分布或lstm_out维度与attention层不匹配。解决检查self.attention输入维度lstm_out.shape[-1]必须等于self.attention.in_features在forward中打印attention_logits.min(), attention_logits.max()若范围 1说明 LSTM 输出太弱调大hidden_dim或加 residual connection临时加attention_logits attention_logits * 10缩放观察是否恢复区分度找到根因后再移除5.4 现象word_tokenize把中文混入的英文词切错如Python代码→[Python, 代, 码]原因word_tokenize是纯英文 tokenizer对中英混杂文本无感知会把中文字符当单字符切分。解决方案 A推荐预处理时用正则分离中英文英文走 NLTK中文走jiebaimport re def hybrid_tokenize(text): # 提取英文片段 en_parts re.findall(r[a-zA-Z], text) # 提取中文片段 zh_parts re.findall(r[\u4e00-\u9fff], text) # 合并保持顺序 tokens [] for char in text: if \u4e00 char \u9fff: tokens.append(char) # 单字切分 or jieba.lcut(char) elif char.isalpha(): # 累积英文单词 pass # 实际项目中建议用 langdetect 先判语种再路由 tokenizer return tokens方案 B改用spacy支持多语种但会增加依赖NLTK 优势丧失5.5 现象模型在 validation 上 acc 95%但线上 inference 全错原因DataLoader的collate_fn在 train/val/inference 时行为不一致——train 用shuffleTrueinference 用shuffleFalse但collate_fn内部逻辑如max_len截断未同步。解决强制统一inference 时也用DataLoaderbatch_size1shuffleFalsecollate_fn完全复用训练版关键检查打印 inference batch 的x.shape和lengths确认与 train 一致终极保险inference 前加model.eval()和torch.no_grad()禁用 dropout/batchnorm6. 进阶技巧如何用 3 行代码把 BiLSTM 升级为 RoBERTa 微调上面的 BiLSTM 是教学 baseline但真实项目往往需要更高精度。好消息是NLTK 预处理 PyTorch 模型的架构无缝兼容 Hugging Face Transformers。你不需要重写整个 pipeline只需替换模型和 tokenizer其余数据加载、loss 计算、eval loop全复用。6.1 替换 tokenizer用AutoTokenizer替代nltk.word_tokenizefrom transformers import AutoTokenizer # 加载 RoBERTa tokenizer自动适配模型 tokenizer AutoTokenizer.from_pretrained(roberta-base) # 替换原来的 NLTK 分词 def tokenize_with_roberta(text): # RoBERTa tokenizer 内置 lower, special tokens, truncation return tokenizer( text, truncationTrue, paddingTrue, max_length128, return_tensorspt ) # 输出{input_ids: tensor([[...]]), attention_mask: tensor([[...]])} # 直接喂给 model(input_ids, attention_mask)逻辑说明RoBERTa tokenizer 比 NLTK 复杂得多——它处理 subwordunacceptable → un accept able、添加[CLS]/[SEP]、生成attention_mask。但你无需关心细节AutoTokenizer自动匹配模型要求。关键是NLTK 的输出是list[str]而 Transformers 的输入是dict[tensor]所以 collate_fn 必须重写见 6.2。6.2 重写 collate_fn支持input_idsattention_mask的 batch 合并def roberta_collate_fn(batch, tokenizer): texts, labels zip(*batch) # tokenizer.batch_encode_plus 是核心 encoded tokenizer( list(texts), truncationTrue, paddingTrue, max_length128, return_tensorspt ) return ( encoded[input_ids], encoded[attention_mask], torch.tensor(labels, dtypetorch.long) ) # DataLoader 实例化仅改 collate_fn dataloader DataLoader( dataset, batch_size16, shuffleTrue, collate_fnlambda b: roberta_collate_fn(b, tokenizer) )逻辑说明tokenizer(..., return_tensorspt)直接返回input_ids和本文还有配套的精品资源点击获取
返回列表