基于PyTorch张量操作实现多维张量按起止值筛选(含补零)
PyTorch张量高效提取起止值间内容的解决方案
需求说明
- 输入:固定长度的二维句子张量(已分词补零),指定起始值
start_value和结束值end_value - 核心规则:
- 对每一行,找到第一个
start_value的索引s_idx和第一个end_value的索引e_idx - 若
s_idx/e_idx不存在,或s_idx >= e_idx,返回与原行长度一致的全零张量 - 否则提取
s_idx+1到e_idx-1之间的元素,其余位置补零,保持输出张量与输入形状完全对称
- 对每一行,找到第一个
示例输入
import torch start_value, end_value = 4,9 data = torch.tensor([ [3,4,7,8,9,2,0,0,0,0], [1,5,3,4,7,2,8,9,10,0], [3,4,7,8,10,0,0,0,0,0], # 不含结束值 [3,7,5,9,2,0,0,0,0,0], # 不含起始值 ])
高效张量操作实现
def extract_between_values(data: torch.Tensor, s_id: int, e_id: int) -> torch.Tensor: batch_size, seq_len = data.shape device = data.device # 批量定位每行第一个起始值的索引,不存在则设为seq_len(超出有效范围) s_mask = (data == s_id) s_idx = torch.argmax(s_mask.int(), dim=1) s_idx = torch.where(s_mask.any(dim=1), s_idx, torch.tensor(seq_len, device=device)) # 批量定位每行第一个结束值的索引,不存在则设为-1(超出有效范围) e_mask = (data == e_id) e_idx = torch.argmax(e_mask.int(), dim=1) e_idx = torch.where(e_mask.any(dim=1), e_idx, torch.tensor(-1, device=device)) # 生成全局位置矩阵,广播对比得到有效区间掩码 pos_indices = torch.arange(seq_len, device=device).repeat(batch_size, 1) valid_mask = (pos_indices > s_idx.unsqueeze(1)) & (pos_indices < e_idx.unsqueeze(1)) # 保留有效区间元素,其余补零 result = torch.where(valid_mask, data, torch.zeros_like(data)) # 修正起止索引无效的行,强制设为全零 invalid_rows = (s_idx >= e_idx) result[invalid_rows] = 0 return result
代码关键逻辑说明
- 批量索引定位:用
argmax一次性获取所有行的目标值索引,避免逐行循环;通过torch.where修正无目标值的情况,确保后续判断逻辑统一 - 有效区间掩码:利用张量广播特性,生成全局位置矩阵与每行起止索引对比,快速得到所有有效位置的掩码
- 对称输出保证:全程保持与输入相同的张量形状,解决原循环代码返回变长列表的问题
- 性能优势:所有操作均为PyTorch底层优化的张量运算,避免Python循环的开销,数据量越大,速度提升越明显
测试输出
调用函数后得到的固定长度对称结果:
result = extract_between_values(data, start_value, end_value) print(result) # 输出: # tensor([[0, 0, 7, 8, 0, 0, 0, 0, 0, 0], # [0, 0, 0, 0, 7, 2, 8, 0, 0, 0], # [0, 0, 0, 0, 0, 0, 0, 0, 0, 0], # [0, 0, 0, 0, 0, 0, 0, 0, 0, 0]])
内容的提问来源于stack exchange,提问作者adib.mosharrof
相关产品推荐
相关产品推荐

