【发布时间】:2021-11-13 06:49:29
【问题描述】:
假设我有一个长度为N 的数字列表lst,以及两个数字epsilon 和tau。我想找到(N,N,N) 掩码矩阵mask 这样mask[i][j][k]=1 当且仅当
abs(lst[i] - lst[j]) <= epsilon and abs(lst[i] - lst[k]) >= tau
这是我尝试过的:
d_mat = torch.cdist(lst.unsqueeze(0), lst.unsqueeze(0))
within_eps = torch.where(dmat <= eps, 1, 0)
over_tau = torch.where(dmat >= tau, 1, 0)
mask = torch.zeros((N,N,N))
for i in range(N):
for j in range(N):
for k in range(N):
if within_eps[i][j] == 1 and over_tau[i][k] == 1:
mask[i][j][k] = 1
else:
mask[i][j][k] = 0
所以基本上我是天真地做到了。您能否通过步骤向我展示您是如何为此提出矢量化的?
【问题讨论】:
标签: python-3.x pytorch vectorization