You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于PyTorch张量操作实现多维张量按起止值筛选(含补零)

PyTorch张量高效提取起止值间内容的解决方案

需求说明

  • 输入:固定长度的二维句子张量(已分词补零),指定起始值start_value和结束值end_value
  • 核心规则:
    1. 对每一行,找到第一个start_value的索引s_idx和第一个end_value的索引e_idx
    2. 若s_idx/e_idx不存在,或s_idx >= e_idx,返回与原行长度一致的全零张量
    3. 否则提取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

代码关键逻辑说明

  1. 批量索引定位:用argmax一次性获取所有行的目标值索引,避免逐行循环;通过torch.where修正无目标值的情况,确保后续判断逻辑统一
  2. 有效区间掩码:利用张量广播特性,生成全局位置矩阵与每行起止索引对比,快速得到所有有效位置的掩码
  3. 对称输出保证:全程保持与输入相同的张量形状,解决原循环代码返回变长列表的问题
  4. 性能优势:所有操作均为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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.15 11:05:22