咨询适用于BERT的文本中间截断工具:单词多Token场景适配
实现BERT文本中间截断的工具库推荐及实现
首选工具:Hugging Face Transformers
Hugging Face的Transformers库是处理BERT相关任务的标准工具,自带BERT系列模型的tokenizer,完全支持自定义截断逻辑,包括你需要的中间截断。
以下是具体的实现代码,考虑了BERT要求的[CLS]和[SEP]特殊token,确保最终输入的总token数不超过512:
from transformers import BertTokenizer # 初始化BERT tokenizer(可替换为你使用的具体BERT模型,如bert-base-chinese) tokenizer = BertTokenizer.from_pretrained('bert-base-uncased') def bert_middle_truncate(text, max_total_tokens=512): # 先将文本转为BERT的子词token,暂不添加特殊token subword_tokens = tokenizer.tokenize(text) total_subwords = len(subword_tokens) # 预留[CLS]和[SEP]的位置,实际可用于文本的token数为max_total_tokens - 2 available_tokens = max_total_tokens - 2 if total_subwords <= available_tokens: # 文本长度符合要求,直接返回标准编码结果 return tokenizer.encode_plus( text, add_special_tokens=True, max_length=max_total_tokens, padding='max_length', truncation=False, return_tensors=None ) # 计算需要截断的token数量 excess_tokens = total_subwords - available_tokens # 从开头和结尾各截断一部分,优先均分,确保中间部分保留 truncate_from_start = excess_tokens // 2 truncate_from_end = excess_tokens - truncate_from_start # 截取中间的子词token truncated_subwords = subword_tokens[truncate_from_start : total_subwords - truncate_from_end] # 添加特殊token并转为模型可识别的输入格式 input_ids = tokenizer.build_inputs_with_special_tokens( tokenizer.convert_tokens_to_ids(truncated_subwords) ) attention_mask = [1] * len(input_ids) # 补全到指定的最大长度 padding_length = max_total_tokens - len(input_ids) input_ids += [tokenizer.pad_token_id] * padding_length attention_mask += [0] * padding_length return { 'input_ids': input_ids, 'attention_mask': attention_mask }
代码说明
- Tokenize处理:先将文本拆分为BERT的子词token,避免直接按单词截断导致的子词拆分问题;
- 特殊token预留:BERT要求输入开头为
[CLS]、结尾为[SEP],因此实际可用于文本的token数为510; - 中间截断逻辑:计算超出的token数后,从文本的开头和结尾各截断一部分,确保核心的中间内容被保留;
- 格式补全:最终生成符合BERT输入要求的
input_ids和attention_mask,并补全到指定长度。
其他可选方式
如果你需要更灵活的文本预处理流程,可以结合spaCy等NLP库先做分句/分段,再用上述逻辑处理,但Transformers库已经能满足绝大多数场景的需求。
内容的提问来源于stack exchange,提问作者Tuan Do
相关产品推荐
相关产品推荐

