如何基于1D布尔张量创建'islands'风格PyTorch矩阵
问题定义
现有长度为N的1维稀疏布尔张量(取值仅为0/1),需要生成尺寸为[N,N]的2维布尔张量,输出张量需呈现输入诱导的「岛屿」结构:当且仅当索引对(i,j)对应的闭区间[min(i,j), max(i,j)]完全落在输入的连续1片段(即岛屿段)内时,输出张量对应位置为1,否则为0。结构对应示意:输入1维张量排列在上方,输出矩阵排列在下方,每个连续1段会对应矩阵中一个贴在对角线位置的正方形全1块。
实现方案
核心思路:给输入中每个连续的同值片段分配唯一ID,将所有0值位置的ID统一设为不可能匹配的特殊值,最后通过广播比较两个位置的片段ID,ID相同则说明两个位置属于同一个1的连续段,对应输出位置为1。
以下给出两种最常用的工程实现,均为无循环的向量化写法,对稀疏输入友好,大尺寸张量下运行效率极高。
NumPy 实现(CPU通用场景)
import numpy as np def build_island_matrix(input_1d: np.ndarray) -> np.ndarray: # 检测相邻位置的数值跳变,累计和生成每个位置的片段ID segment_id = np.cumsum( np.concatenate([[0], (input_1d[:-1] != input_1d[1:]).astype(int)]) ) # 所有0值位置的ID设为-1,避免跨片段误匹配 segment_id[input_1d == 0] = -1 # 广播比较两两位置的片段ID,得到结果 return segment_id[:, None] == segment_id[None, :]
PyTorch 实现(GPU/深度学习工作流场景)
import torch def build_island_matrix(input_1d: torch.Tensor) -> torch.Tensor: # 检测相邻位置的数值跳变,累计和生成每个位置的片段ID segment_id = torch.cumsum( torch.cat([ torch.tensor([0], device=input_1d.device, dtype=torch.long), (input_1d[:-1] != input_1d[1:]).long() ]), dim=0 ) # 所有0值位置的ID设为-1,避免跨片段误匹配 segment_id[input_1d == 0] = -1 # 广播比较两两位置的片段ID,得到结果 return segment_id.unsqueeze(1) == segment_id.unsqueeze(0)
效果验证
以输入[1,1,0,1,0,1,1,1](长度N=8)为例,输出矩阵符合预期:
- 前2行前2列全为1,对应第一个长度为2的连续1段
- 第4行第4列(0基索引为3)单独为1,对应第二个长度为1的连续1段
- 最后3行最后3列全为1,对应第三个长度为3的连续1段
- 其余位置全为0,无跨段的误匹配。
内容的提问来源于stack exchange,提问作者Codevan
相关产品推荐
相关产品推荐

