- html - 出于某种原因,IE8 对我的 Sass 文件中继承的 html5 CSS 不友好?
- JMeter 在响应断言中使用 span 标签的问题
- html - 在 :hover and :active? 上具有不同效果的 CSS 动画
- html - 相对于居中的 html 内容固定的 CSS 重复背景?
我有一批形状为 (bs, m, n) 的向量(即维度为 mxn 的 bs 向量)。对于每个批处理,我想计算第一个向量与其余 (m-1) 个向量的 Jaccard 相似度
例子:
a = [
[[3, 8, 6, 8, 7],
[9, 7, 4, 8, 1],
[7, 8, 8, 5, 7],
[3, 9, 9, 4, 4]],
[[7, 3, 8, 1, 7],
[3, 0, 3, 4, 2],
[9, 1, 6, 1, 6],
[2, 7, 0, 6, 6]]
]
找出 a[:,0,:] 和 a[:,1:,:] 之间的成对 jaccard 相似度即,
[3, 8, 6, 8, 7] with each of [[9, 7, 4, 8, 1], [7, 8, 8, 5, 7], [3, 9, 9, 4, 4]] (3 scores)
and
[7, 3, 8, 1, 7] with each of [[3, 0, 3, 4, 2], [9, 1, 6, 1, 6], [2, 7, 0, 6, 6]] (3 scores)
这是我试过的Jaccard函数
def js(la1, la2):
combined = torch.cat((la1, la2))
union, counts = combined.unique(return_counts=True)
intersection = union[counts > 1]
torch.numel(intersection) / torch.numel(union)
虽然这适用于大小不等的张量,但这种方法的问题是每个组合(张量对)中唯一值的数量可能不同,并且由于 PyTorch 不支持锯齿状张量,我无法处理一次批量处理向量。
如果我无法以预期的清晰度表达问题,请告诉我。在这方面的任何帮助将不胜感激
编辑:这是通过迭代第一维和第二维实现的流程。我希望有一个用于批处理的以下代码的矢量化版本
bs = 2
m = 4
n = 5
a = torch.randint(0, 10, (bs, m, n))
print(f"Array is: \n{a}")
for bs_idx in range(bs):
first = a[bs_idx,0,:]
for row in range(1, m):
second = a[bs_idx,row,:]
idx = js(first, second)
print(f'comparing{first} and {second}: {idx}')
最佳答案
我不知道如何在 pytorch 中实现这一点,因为 AFAIK pytorch 不支持张量上的集合操作。在您的 js()
实现中,union
计算应该有效,但是 intersection = union[counts > 1]
没有给您正确的结果如果其中一个张量包含重复值。另一方面,Numpy 内置了对 union1d
和 intersect1d
的支持。您可以使用 numpy 向量化来计算成对的 jaccard 索引,而无需使用 for 循环:
import numpy as np
def num_intersection(vec1: np.ndarray, vec2: np.ndarray) -> int:
return np.intersect1d(vec1, vec2, assume_unique=False).size
def num_union(vec1: np.ndarray, vec2: np.ndarray) -> int:
return np.union1d(vec1, vec2).size
def jaccard1d(vec1: np.ndarray, vec2: np.ndarray) -> float:
assert vec1.ndim == vec2.ndim == 1 and vec1.shape[0] == vec2.shape[0], 'vec1 and vec2 must be 1D arrays of equal length'
return num_intersection(vec1, vec2) / num_union(vec1, vec2)
jaccard2d = np.vectorize(jaccard1d, signature='(m),(n)->()')
def jaccard(vecs1: np.ndarray, vecs2: np.ndarray) -> np.ndarray:
"""
Return intersection-over-union (Jaccard index) between two sets of vectors.
Both sets of vectors are expected to be flattened to 2D, where dim 0 is the batch
dimension and dim 1 contains the flattened vectors of length V (jaccard index of
an n-dimensional vector and of its flattened 1D-vector is equal).
Args:
vecs1 (ndarray[N, V]): first set of vectors
vecs2 (ndarray[M, V]): second set of vectors
Returns:
ndarray[N, M]: the NxM matrix containing the pairwise jaccard indices for every vector in vecs1 and vecs2
"""
assert vecs1.ndim == vecs2.ndim == 2 and vecs1.shape[1] == vecs2.shape[1], 'vecs1 and vecs2 must be 2D arrays with equal length in axis 1'
return jaccard2d(vecs1, vecs2)
这当然不是最优的,因为代码不在 GPU 上运行。如果我使用形状 (1, 10)
的 vecs1
和形状 的
我在我的机器上得到的平均循环时间为 vecs2
运行 jaccard
函数(10_000, 10)200 ms ± 1.34 ms
,这对于大多数用例来说应该足够快了。 pytorch 和 numpy 数组之间的转换非常便宜。
要将此函数应用于数组 a
的问题:
a = torch.tensor(a).numpy() # just to demonstrate
ious = [jaccard(batch[:1, :], batch[1:, :]) for batch in a]
np.array(ious).squeeze() # 2 batches with 3 scores each -> 2x3 matrix
# array([[0.28571429, 0.4 , 0.16666667],
# [0.14285714, 0.16666667, 0.14285714]])
如果需要,对结果使用 torch.from_numpy()
再次获得 pytorch 张量。
如果你需要一个pytorch版本来计算Jaccard索引,我在torch中部分实现了numpy的intersect1d
:
from torch import Tensor
def torch_intersect1d(t1: Tensor, t2: Tensor, assume_unique: bool = False) -> Tensor:
if t1.ndim > 1:
t1 = t1.flatten()
if t2.ndim > 1:
t2 = t2.flatten()
if not assume_unique:
t1 = t1.unique(sorted=True)
t2 = t2.unique(sorted=True)
# generate a m x n intersection matrix where m is numel(t1) and n is numel(t2)
intersect = t1[(t1.view(-1, 1) == t2.view(1, -1)).any(dim=1)]
if not assume_unique:
intersect = intersect.sort().values
return intersect
def torch_union1d(t1: Tensor, t2: Tensor) -> Tensor:
return torch.cat((t1.flatten(), t2.flatten())).unique()
def torch_jaccard1d(t1: Tensor, t2: Tensor) -> float:
return torch_intersect1d(t1, t2).numel() / torch_union1d(t1, t2).numel()
要向量化 torch_jaccard1d
函数,您可能需要查看 torch.vmap
,它允许您在任意批处理维度上对函数进行矢量化(类似于 numpy 的 vectorize
)。 vmap
函数是一个原型(prototype)功能,在通常的 pytorch 发行版中尚不可用,但您可以使用每晚构建的 pytorch 来获取它。我还没有测试过,但这可能有效。
关于python - 在 PyTorch 中找到一批向量之间的 jaccard 相似性,我们在Stack Overflow上找到一个类似的问题: https://stackoverflow.com/questions/72657212/
我需要修复 getLineNumberFor 方法,以便如果 lastName 的第一个字符位于 A 和 M 之间,则返回 1;如果它位于 N 和 Z 之间,则返回 2。 在我看来听起来很简单,但我不
您好,感谢您的帮助!我有这个: 0 我必须在每次点击后增加“pinli
Javascript 中是否有一种方法可以在不使用 if 语句的情况下通过 switch case 结构将一个整数与另一个整数进行比较? 例如。 switch(integer) { case
我有一列是“日期”类型的。如何在自定义选项中使用“之间”选项? 最佳答案 请注意,您有2个盒子。 between(在SQL中)包含所有内容,因此将框1设置为:DATE >= startdate,将框2
我有一个表,其中包含年、月和一些数字列 Year Month Total 2011 10 100 2011 11 150 2011 12 100 20
这个问题已经有答案了: Extract a substring between double quotes with regular expression in Java (2 个回答) how to
我有一个带有类别的边栏。正如你在这里看到的:http://kees.een-site-bouwen.nl/ url 中类别的 ID。带有 uri 段(3)当您单击其中一个类别时,例如网页设计。显示了一
这个问题在这里已经有了答案: My regex is matching too much. How do I make it stop? [duplicate] (5 个答案) 关闭 4 年前。 我
我很不会写正则表达式。 我正在尝试获取括号“()”之间的值。像下面这样的东西...... $a = "POLYGON((1 1,2 2,3 3,1 1))"; preg_match_all("/\((
我必须添加一个叠加层 (ImageView),以便它稍微移动到包含布局的左边界的左侧。 执行此操作的最佳方法是什么? 尝试了一些简单的方法,比如将 ImageView 放在布局中并使用负边距 andr
Rx 中是否有一些扩展方法来完成下面的场景? 我有一个开始泵送的值(绿色圆圈)和其他停止泵送的值(簧片圆圈),蓝色圆圈应该是预期值,我不希望这个命令被取消并重新创建(即“TakeUntil”和“Ski
我有一个看起来像这样的数据框(Dataframe X): id number found 1 5225 NA 2 2222 NA 3 3121 NA 我有另一个看起来
所以,我正在尝试制作正则表达式,它将解析存储在对象中的所有全局函数声明,例如,像这样 const a = () => {} 我做了这样的事情: /(?:const|let|var)\s*([A-z0-
我正在尝试从 Intellivision 重新创建 Astro-Smash,我想让桶保持在两个 Angular 之间。我只是想不出在哪里以及如何让这个东西停留在两者之间。 我已经以各种方式交换了函数,
到处检查但找不到答案。 我有这个页面,我使用 INNER JOIN 将两个表连接在一起,获取它们的值并显示它们。我有这个表格,用来获取变量(例如开始日期、结束日期和卡号),这些变量将作为从表中调用值的
我陷入了两个不同的问题/错误之间,无法想出一个合适的解决方案。任何帮助将不胜感激 上下文、FFI 和调用大量 C 函数,并将 C 类型包装在 rust 结构中。 第一个问题是ICE: this pat
我在 MySQL 中有一个用户列表,在订阅时,时间戳是使用 CURRENT_TIMESTAMP 在数据库中设置的。 现在我想从此表中选择订阅日期介于第 X 天和第 Y 天之间的表我尝试了几个查询,但不
我的输入是开始日期和结束日期。我想检查它是在 12 月 1 日到 3 月 31 日之间。(年份可以更改,并且只有在此期间内或之外的日期)。 到目前为止,我还没有找到任何关于 Joda-time 的解决
我正在努力了解线程与 CPU 使用率的关系。有很多关于线程与多处理的讨论(一个很好的概述是 this answer )所以我决定通过在运行 Windows 10、Python 3.4 的 8 CPU
我正在尝试编写 PHP 代码来循环遍历数组以创建 HTML 表格。我一直在尝试做类似的事情: fetchAll(PDO::FETCH_ASSOC); ?>
我是一名优秀的程序员,十分优秀!