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

资讯详情

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

【Bug已解决】Allow TextClassificationPipeline to handle input longer than model_max_length tokens 解决方案

【Bug已解决】Allow TextClassificationPipeline to handle input longer than model_max_length tokens 解决方案 【Bug已解决】Allow TextClassificationPipeline to handle input longer than model_max_length tokens 解决方案一、现象长什么样用text-classificationpipeline 对一段超过模型model_max_length的文本做分类时行为不符合预期# 现象 A超长输入被静默截断只看了开头 # 输入 5000 token模型 max_length512pipeline 只喂了前 512 token 给模型 # 剩下的 4488 token 直接丢弃 - 分类结果只反映开头漏掉关键尾部信息 # 现象 B传 truncationFalse 直接报错 ValueError: token sequence length (5000) exceeds model maximum length (512) # 关闭截断后模型 forward 因序列超长而报错 # 现象 C直接 OOM # 有人为绕过截断把整段塞进去序列 5000 远超 512显存爆炸 # 典型触发 from transformers import pipeline pipe pipeline(text-classification, modelbert-base-uncased) out pipe(very_long_text) # 只分类了前 512 token尾部被忽略最典型的指纹短文本分类正常长文档分类看起来正常但结果不对——因为 pipeline 默认把超长输入截断到model_max_length用户并不知道后半段被丢了。二、背景TextClassificationPipeline的默认行为是为每个输入构造一个不超过model_max_length的序列通过truncationTrue然后整体过一次模型。这对短文本没问题但对长文档论文、合同、长评论就有两个问题信息丢失截断只保留前model_max_lengthtoken关键内容可能在尾部如综上所述本文结论是...。无法关闭截断设truncationFalse又会让模型 forward 因超长报错。正确的做法也是这个 issue 想要的是对超长输入做分块chunk 步长stride窗口滑动每个窗口单独分类再把各窗口的结果聚合如取平均概率 / 投票从而在不超长的前提下利用全文信息。这是长文本分类的标准范式。三、根因根因有两类pipeline 默认整体截断无分块聚合逻辑。TextClassificationPipeline的预处理把整段文本 tokenize 后按max_length截断成单条没有切成多个重叠窗口分别推理再聚合的代码路径。于是超长输入要么丢信息默认截断要么报错关截断。没有窗口级推理 聚合的接口。 即使想手动分块pipeline 也没有参数让你指定max_chunk_length/stride/aggregation用户只能自己写后处理体验割裂。四、最小可运行复现下面用纯 Python 模拟超长输入被整体截断只取前 N 个 token与分块聚合的差异from typing import List def classify_chunk(tokens: List[str]) - float: 模拟某窗口的分类得分这里用窗口里 positive 词出现比例示意。 if not tokens: return 0.0 pos sum(1 for t in tokens if t good) return pos / len(tokens) def pipeline_default_truncate(tokens: List[str], max_len: int) - float: 有 bug整体截断到 max_len只看前 max_len 个 token。 return classify_chunk(tokens[:max_len]) def pipeline_chunked(tokens: List[str], max_len: int, stride: int) - float: 修正滑动窗口分块各块得分取平均。 if len(tokens) max_len: return classify_chunk(tokens) scores [] start 0 while start len(tokens): chunk tokens[start:start max_len] scores.append(classify_chunk(chunk)) if start max_len len(tokens): break start stride return sum(scores) / len(scores) # 文档开头全是中性词尾部有 positive 信号 doc [neutral]*600 [good, good, good] max_len, stride 512, 400 default pipeline_default_truncate(doc, max_len) chunked pipeline_chunked(doc, max_len, stride) print(默认截断得分:, round(default, 4)) # 0.0没看到尾部 good print(分块聚合得分:, round(chunked, 4)) # 0看到尾部信号 assert chunked default, 复现失败默认截断丢失了尾部信息运行后默认截断只看了前 512 个 neutral 词得分 0分块聚合则通过滑动窗口看到了尾部 3 个 good得分 0复现并修复了根因 1。五、解决方案第一层最小直接修复最快的止血在调用 pipeline 前手动把长文本切成重叠窗口分别推理再聚合得分from transformers import pipeline, AutoTokenizer pipe pipeline(text-classification, modelbert-base-uncased, tokenizerAutoTokenizer.from_pretrained(bert-base-uncased)) tok pipe.tokenizer max_len tok.model_max_length # 如 512 stride 400 def classify_long(text: str, aggmean): 第一层修复长文本滑动窗口分块 聚合。 enc tok(text, add_special_tokensFalse)[input_ids] if len(enc) max_len: return pipe(text)[0] scores [] start 0 while start len(enc): chunk_ids enc[start:start max_len] chunk_text tok.decode(chunk_ids) res pipe(chunk_text)[0] scores.append(res[score] if res[label] POSITIVE else 1 - res[score]) if start max_len len(enc): break start stride agg_score sum(scores) / len(scores) if agg mean else max(scores) return {label: POSITIVE if agg_score 0.5 else NEGATIVE, score: agg_score} out classify_long(very_long_text)第一层让用户立刻能对超长文本做全文感知的分类不再只看开头。六、解决方案第二层结构性改进用LongTextClassifier把分块 stride 聚合封装进 pipeline 扩展支持参数化from dataclasses import dataclass from typing import Callable, List dataclass class LongTextClassifier: 支持超长输入的文本分类滑动窗口分块 聚合。 pipe: Callable max_len: int 512 stride: int 400 agg: str mean # mean / max / vote def _chunk_ids(self, input_ids: List[int]) - List[List[int]]: if len(input_ids) self.max_len: return [input_ids] chunks [] s 0 while s len(input_ids): chunks.append(input_ids[s:s self.max_len]) if s self.max_len len(input_ids): break s self.stride return chunks def classify(self, text: str, tokenizer) - dict: ids tokenizer(text, add_special_tokensFalse)[input_ids] chunks self._chunk_ids(ids) pos_scores [] for c in chunks: t tokenizer.decode(c) r self.pipe(t)[0] pos_scores.append(r[score] if r[label] POSITIVE else 1 - r[score]) if self.agg mean: score sum(pos_scores) / len(pos_scores) elif self.agg max: score max(pos_scores) else: # vote score sum(1 for s in pos_scores if s 0.5) / len(pos_scores) return {label: POSITIVE if score 0.5 else NEGATIVE, score: score} # 使用 cls LongTextClassifier(pipepipe, max_len512, stride400, aggmean) print(cls.classify(very_long_text, tok))LongTextClassifier把长文本分类的窗口逻辑参数化用户只需指定max_len/stride/agg不必每次手写分块也方便处理尾部关键词类任务。七、解决方案第三层断言 / CI 守护用 pytest 固化超长输入被分块且不丢尾部信息、各块不超 max_lenimport pytest def test_long_input_chunked_not_truncated(): from long_text_cls import LongTextClassifier calls [] def fake_pipe(t): calls.append(t); return [{label: POSITIVE, score: 0.9}] cls LongTextClassifier(pipefake_pipe, max_len8, stride6, aggmean) # 模拟 tokenizer每个字符一个 id class T: def __call__(self, t, add_special_tokensFalse): return {input_ids: list(range(len(t)))} def decode(self, ids): return .join(str(i) for i in ids) res cls.classify(x*20, T()) # 20 8应分块 # 至少被切成多块 assert len(calls) 2, 超长输入应被分块而非整体截断 # 每块 token 数不超过 max_len assert all(len(c) 8 for c in calls) def test_aggregation_sees_tail(): from long_text_cls import LongTextClassifier # 前块负、后块正mean 应介于之间vote 应偏正 def fake_pipe(t): # 用块内容判断含 P 的块判正 return [{label: POSITIVE if P in t else NEGATIVE, score: 0.9 if P in t else 0.1}] class T: def __call__(self, t, add_special_tokensFalse): return {input_ids: list(t)} def decode(self, ids): return .join(ids) cls LongTextClassifier(pipefake_pipe, max_len5, stride4, aggmean) res cls.classify(NNNNNPPP, T()) assert res[label] POSITIVE, 分块聚合应看到尾部正信号CI 跑pytest tests/test_long_text_classification.py以后只要有人又让 pipeline 对超长输入静默截断丢尾部测试立刻红灯。八、排查清单当text-classification对长文本结果不对按顺序查结果只反映开头 → 默认整体截断到model_max_length尾部被丢用分块聚合。关truncationFalse报超长 → 别关截断改用分块每块不超 max_len。直接塞整段 OOM → 必须分块单块不超 max_len。聚合策略关键词在尾部用mean/vote整体情感用mean要抓任一处正信号用max。长期方案用LongTextClassifier把分块stride聚合做成标准能力避免每次手写。九、小结Allow TextClassificationPipeline to handle input longer than model_max_length 的根因是pipeline 默认把超长输入整体截断到model_max_length只取前 N 个 token 过模型尾部信息丢失关闭截断又会让模型因超长报错于是长文档分类看起来正常结果却错。第一层手动滑动窗口分块 聚合得分立刻实现全文感知分类。第二层用LongTextClassifier把分块/stride/聚合参数化用户指定窗口即可不必手写。第三层pytest 断言超长输入被分块、每块不超 max_len、聚合看到尾部信号防止回归。记住模型有model_max_length上限长文本分类不能靠整体截断蒙混滑动窗口分块 聚合mean/max/vote才是利用全文的标准做法。
返回列表