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

如何使用PyTorch或Numpy在布尔张量中查找2D模式并更新频次字典

PyTorch实现2D布尔模式匹配与频次更新

核心思路

完全基于PyTorch原生算子实现,无手动编写的Python层低效循环,核心利用滑动窗口展开+批量广播比较实现多模式并行匹配:

  • 用torch.nn.functional.unfold算子从大尺寸输入张量中提取所有和目标模式尺寸一致的滑动窗口
  • 所有模式批量和滑动窗口做全等比较,直接统计匹配数更新字典

PyTorch完整实现

import torch
import torch.nn.functional as F

def update_pattern_frequency(input_tensor: torch.Tensor, pattern_dict: dict, patch_size: int = 6) -> None:
    """
    匹配输入张量中的2D布尔模式,自动更新模式字典的频次
    Args:
        input_tensor: 输入大尺寸布尔张量,维度为[通道数, 高度, 宽度],例如3x256x256
        pattern_dict: 存储2D布尔模式的字典,键为patch_size x patch_size的布尔张量,值为对应出现频次
        patch_size: 2D模式的尺寸,默认为6
    """
    # 输入维度校验
    assert input_tensor.ndim == 3, "输入张量维度需为[C, H, W]"
    device = input_tensor.device

    # 1. 提取输入张量的所有滑动窗口,展开为二维矩阵
    # 输出形状:[窗口总数量, 通道数*patch_size*patch_size]
    unfolded = F.unfold(
        input_tensor.unsqueeze(0), # 补充batch维度
        kernel_size=patch_size,
        stride=1 # 步长为1表示匹配所有重叠窗口,若不需要重叠可设为patch_size
    ).squeeze(0).T

    # 2. 把所有模式展平并拼接为批量矩阵
    patterns = torch.stack([p.flatten() for p in pattern_dict.keys()]).to(device)

    # 3. 批量广播比较,找出完全匹配的窗口
    # 所有像素值完全相等才判定为匹配
    match_mask = (unfolded.unsqueeze(1) == patterns.unsqueeze(0)).all(dim=-1)

    # 4. 统计每个模式的匹配数量,更新字典
    match_counts = match_mask.sum(dim=0).cpu().tolist()
    for idx, pattern in enumerate(pattern_dict.keys()):
        pattern_dict[pattern] += match_counts[idx]

使用示例

# 模拟输入:3通道256x256布尔张量
input_tensor = torch.randint(0, 2, (3, 256, 256), dtype=torch.bool)
# 模拟模式字典:3个6x6布尔模式,初始频次为0
pattern_dict = {
    torch.randint(0, 2, (3, 6, 6), dtype=torch.bool): 0,
    torch.randint(0, 2, (3, 6, 6), dtype=torch.bool): 0,
    torch.randint(0, 2, (3, 6, 6), dtype=torch.bool): 0
}
# 执行匹配更新
update_pattern_frequency(input_tensor, pattern_dict, patch_size=6)
# 输出更新后的频次
print(pattern_dict.values())

可选Numpy实现

逻辑和PyTorch版本完全一致,使用Numpy原生滑动窗口算子实现:

import numpy as np

def update_pattern_frequency_np(input_array: np.ndarray, pattern_dict: dict, patch_size: int = 6) -> None:
    # 提取所有滑动窗口
    windows = np.lib.stride_tricks.sliding_window_view(input_array, (patch_size, patch_size), axis=(1,2))
    windows = windows.reshape(-1, input_array.shape[0]*patch_size*patch_size)
    # 批量比较
    patterns = np.stack([p.flatten() for p in pattern_dict.keys()])
    match_mask = (windows[:, None] == patterns[None, :]).all(axis=-1)
    # 更新字典
    match_counts = match_mask.sum(axis=0)
    for idx, pattern in enumerate(pattern_dict.keys()):
        pattern_dict[pattern] += match_counts[idx]

注意事项

  • 如果需要匹配的是单通道2D模式,可将输入张量按通道拆分后分别调用上述函数
  • 匹配前需保证输入张量和模式的dtype完全一致,避免因uint8和bool类型差异导致匹配失败
  • 所有核心运算均为框架原生优化实现,支持GPU加速,处理大尺寸张量时效率远高于手动编写的Python循环

内容的提问来源于stack exchange,提问作者Yousef

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 05:24:00