PyTorch中获取N×N矩阵内过目标点与中心原点的直线元素的方法咨询
PyTorch中获取N×N矩阵内过目标点与中心原点的直线元素的方法咨询
嗨,我完全get到你的需求了!你现在有个N×N的PyTorch矩阵,把矩阵中心当作原点,给定一个目标点(i,j),想要提取所有经过这个点和中心的直线上的元素,之前试了torch.diag但发现它没法对准矩阵中心——确实,这个函数是针对矩阵角落出发的对角线设计的,完全不适合你的场景。
先看你给出的代码片段,你已经在构建kx/ky网格,还对左半部分的k_grid取反了,接下来咱们把这个功能补全:
首先得明确几个关键点:
- 矩阵的中心坐标:对于N×N矩阵,咱们取
(n//2, n//2)作为中心(n是矩阵边长,不管奇偶都适用,对应你代码里的len(k_grid)//2) - 要找到过中心和目标点的直线上的所有整数坐标点,得先计算目标点相对于中心的偏移量,再通过最大公约数(GCD)确定步长,这样能遍历直线上所有不重复的点
下面是修改并补全后的函数实现:
import torch import math def directionalK(kx, ky, indices): '''Function that provides the K values at a given direction dictated by the indices''' kx_grid, ky_grid = torch.meshgrid(kx, ky, indexing='ij') k_grid = torch.sqrt(kx_grid**2 + ky_grid**2) k_grid[..., :len(k_grid)//2] *= -1 y, x = indices n = k_grid.shape[0] center = (n//2, n//2) # 处理目标点就是中心的特殊情况 if y == center[0] and x == center[1]: return k_grid[center[0], center[1]] # 计算目标点相对中心的偏移量 dy = y - center[0] dx = x - center[1] # 用最大公约数确定最小步长,保证遍历直线上所有整数点 g = math.gcd(abs(dx), abs(dy)) step_y = dy // g step_x = dx // g # 收集直线上所有在矩阵范围内的点 line_points = [] # 先遍历从中心到目标点的方向 curr_y, curr_x = center while 0 <= curr_y < n and 0 <= curr_x < n: line_points.append((curr_y, curr_x)) curr_y += step_y curr_x += step_x # 再遍历中心的反方向(跳过已经加入的中心) curr_y, curr_x = center[0] - step_y, center[1] - step_x while 0 <= curr_y < n and 0 <= curr_x < n: line_points.append((curr_y, curr_x)) curr_y -= step_y curr_x -= step_x # 提取这些点对应的k_grid元素 y_indices = torch.tensor([p[0] for p in line_points], dtype=torch.long) x_indices = torch.tensor([p[1] for p in line_points], dtype=torch.long) line_elements = k_grid[y_indices, x_indices] return line_elements
代码说明:
- 特殊情况处理:如果目标点就是矩阵中心,直接返回中心元素
- 步长计算:通过GCD把偏移量简化为最小整数步长,这样不会跳过直线上的任何整数坐标点,也不会重复遍历
- 双向遍历:从中心出发,分别向目标点方向和反方向遍历,收集所有在矩阵边界内的点
- 索引提取:把收集到的坐标转成PyTorch能识别的长整型张量,直接索引提取对应元素
另外补充下:如果你的N是偶数,(n//2, n//2)是矩阵右下侧的“中心”点,要是你需要的是偶数矩阵的几何中心(介于四个点之间),那得调整中心的定义,但从你的代码逻辑来看,这个实现完全匹配你的需求~
备注:内容来源于stack exchange,提问作者Rotacional
相关产品推荐
相关产品推荐

