PyTorch中如何对批量张量的指定子矩阵应用自定义对角线条带函数
批量处理PyTorch张量提取对角线条带的解决方案
我来帮你搞定这个批量提取对角线条带的问题!首先咱们先把自定义的stripe函数补全并优化成向量化版本(避免循环,效率更高),这个函数专门处理单个i×j(i≥j)的矩阵,提取所有符合要求的对角线条带:
import torch def stripe(a): i, j = a.size() assert i >= j, "矩阵的行数必须大于等于列数" stripe_length = i - j + 1 # 构建行索引:每个条带的行索引是 k + 0到stripe_length-1,k从0到j-1 row_indices = torch.arange(stripe_length, device=a.device)[None, :] + torch.arange(j, device=a.device)[:, None] # 构建列索引:每个条带对应固定的列k,重复stripe_length次 col_indices = torch.arange(j, device=a.device)[:, None].repeat(1, stripe_length) # 提取对应位置的元素,输出形状为 (j, stripe_length) return a[row_indices, col_indices]
接下来处理批量张量的情况,你的输入是形状为[150, 182, 91]的张量,第一维是批量大小,这里有三种常用的处理方式:
方法一:用torch.vmap(推荐,简洁高效)
PyTorch 1.12及以上版本支持vmap,它可以自动将处理单个样本的函数映射到整个批量维度上,不需要手动写循环:
# 示例输入张量 batch_tensor = torch.randn(150, 182, 91) # 用vmap批量应用stripe函数 batch_stripes = torch.vmap(stripe)(batch_tensor) # 输出形状为 [150, 91, 92],对应每个批量样本的91条长度为92的对角线条带
方法二:手动遍历批量(兼容旧版本PyTorch)
如果你用的PyTorch版本不支持vmap,可以手动遍历每个批量样本,处理后再堆叠起来:
batch_stripes_list = [] for single_matrix in batch_tensor: single_stripes = stripe(single_matrix) batch_stripes_list.append(single_stripes) # 将列表堆叠成张量 batch_stripes = torch.stack(batch_stripes_list)
方法三:全向量化操作(极致性能)
如果追求最高的性能,可以直接构建整个批量的索引,一次性提取所有条带,完全避免循环:
batch_size, i, j = batch_tensor.shape stripe_length = i - j + 1 # 构建批量维度的索引,形状为 [batch_size, j, stripe_length] batch_indices = torch.arange(batch_size, device=batch_tensor.device)[:, None, None].repeat(1, j, stripe_length) # 构建行索引,形状为 [batch_size, j, stripe_length] row_indices = (torch.arange(stripe_length, device=batch_tensor.device)[None, :] + torch.arange(j, device=batch_tensor.device)[:, None])[None, :, :].repeat(batch_size, 1, 1) # 构建列索引,形状为 [batch_size, j, stripe_length] col_indices = torch.arange(j, device=batch_tensor.device)[:, None].repeat(1, stripe_length)[None, :, :].repeat(batch_size, 1, 1) # 一次性提取所有批量样本的条带 batch_stripes = batch_tensor[batch_indices, row_indices, col_indices]
这样处理后,batch_stripes的每个元素就是对应批量样本中提取的对角线条带啦~
内容的提问来源于stack exchange,提问作者Ivan Bilan
相关产品推荐
相关产品推荐

