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

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

代码说明:

  1. 特殊情况处理:如果目标点就是矩阵中心,直接返回中心元素
  2. 步长计算:通过GCD把偏移量简化为最小整数步长,这样不会跳过直线上的任何整数坐标点,也不会重复遍历
  3. 双向遍历:从中心出发,分别向目标点方向和反方向遍历,收集所有在矩阵边界内的点
  4. 索引提取:把收集到的坐标转成PyTorch能识别的长整型张量,直接索引提取对应元素

另外补充下:如果你的N是偶数,(n//2, n//2)是矩阵右下侧的“中心”点,要是你需要的是偶数矩阵的几何中心(介于四个点之间),那得调整中心的定义,但从你的代码逻辑来看,这个实现完全匹配你的需求~

备注:内容来源于stack exchange,提问作者Rotacional

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.16 10:34:38