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

如何基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 14:54:36