如何基于PyTorch按句末分割长文本分词张量并限制最大尺寸
解决方案
要实现仅在句末token(1012)处分割张量,同时保证每个子张量长度不超过最大限制(如510),可以按照以下逻辑实现:
核心思路
- 定位所有句末token的位置
- 从左到右遍历,尽可能在不超过最大长度的前提下,累积到最近的句末位置再分割
- 处理剩余的最后一段文本,确保不遗漏
代码实现
import torch # 示例输入张量(二维) tensor = torch.tensor([[3746, 3120, 1024, 1012, 2655, 24051, 2015, 1012]]) # 仅为示例,实际长度14714 tensor = tensor.squeeze(0) # 转换为一维张量,方便处理 max_segment_len = 510 eos_token = 1012 # 获取所有句末token的索引,同时添加张量最后一个位置作为兜底 eos_indices = torch.where(tensor == eos_token)[0] eos_indices = torch.cat([eos_indices, torch.tensor([len(tensor)-1], device=tensor.device)]) segments = [] current_start = 0 for eos_idx in eos_indices: current_segment_length = eos_idx - current_start + 1 # 如果当前句末超出最大长度,取之前最后一个合法的句末分割 if current_segment_length > max_segment_len: # 找到在当前起始位置+最大长度范围内的最后一个句末 valid_eos = eos_indices[eos_indices < current_start + max_segment_len][-1] segments.append(tensor[current_start:valid_eos+1]) current_start = valid_eos + 1 # 如果当前句末未超出最大长度,检查下一个句末是否能纳入,不能则分割 elif current_segment_length <= max_segment_len: # 获取下一个句末的位置 next_eos_candidates = eos_indices[eos_indices > eos_idx] if len(next_eos_candidates) > 0 and (next_eos_candidates[0] - current_start + 1) <= max_segment_len: continue # 无法纳入下一个句末,分割当前片段 segments.append(tensor[current_start:eos_idx+1]) current_start = eos_idx + 1 # 所有token处理完毕则退出循环 if current_start >= len(tensor): break
关键说明
- 先将输入的二维张量转为一维,简化索引操作
- 添加张量最后一个位置到句末索引列表,避免最后一段无句末token时无法分割
- 遍历过程中动态判断是否可以继续累积文本,确保每个子张量长度不超过限制,且仅在句末分割
内容的提问来源于stack exchange,提问作者Alex
相关产品推荐
相关产品推荐

