如何使用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
相关产品推荐
相关产品推荐

