【问题标题】:Binary mask of top n-th quantile in a batch of 2D tensors, but with individual n for each tensor一批 2D 张量中前 n 个分位数的二进制掩码,但每个张量都有单独的 n
【发布时间】:2021-06-27 19:07:38
【问题描述】:

我有一个形状为 (100, 16, 16) 的张量 A 和一个形状为 (100) 的张量 B,其中 100 是批量大小。我想创建一个形状为(100, 16, 16) 的A 的二进制掩码,其中在掩码的每个元素(元素的形状为(1, 16, 16))中,如果该元素大于计算的分位数值,则值为1,否则0。张量 B 中的每个元素依次表示 A 中每个单独元素的百分位值。如果 B 只是一个标量,我可以使用:

flat_A = torch.reshape(A, (100, -1))
quants = torch.quantile(flat_A, B, dim=1)
quants = torch.reshape(quants, (100, 1, 1))
mask = torch.where(A >= quants, 1, 0)
# quants will have shape (100, 1, 1)

问题是:如果 B 是形状为 (100) 的一维张量,就像我上面所说的,我如何计算 A 中每个单独元素的百分位值?我尝试了以下方法,但结果看起来不像我预期的那样:

>>> torch.quantile(flat_A, B, dim=1).shape
torch.Size([100, 100])
>>> torch.quantile(flat_A, B, dim=0).shape
torch.Size([100, 256])

我认为结果的形状应该是(100),所以我可以使用mask = torch.where(A >= quants, 1, 0),还是我误解了?

为了更多的上下文,这个问题也是我之前here的标量B值问题的扩展。

【问题讨论】:

    标签: python pytorch


    【解决方案1】:

    这是使用torch.quantile() 函数的一种方式。请注意,为了简单起见,这里我使用形状为 (5, 2, 2) 而不是 (100, 16, 16) 的张量。

    import torch
    # Generate some data of shape (5, 2, 2)
    A = torch.arange(5 * 2 * 2).reshape(5, 2, 2) + 1.0
    B = torch.linspace(0, 1, 5) # 5 quantile values for each element in A
    
    Af = A.reshape(A.shape[0], -1) # flattens A to a 2D tensor
    quantiles = torch.quantile(Af, B, dim = 1, keepdim = True)
    quants = quantiles[torch.arange(A.shape[0]), torch.arange(A.shape[0]), 0]
    
    mask = (A >= quants[:, None, None]).type(torch.uint8)
    

    这里的张量quantiles 的形状为torch.Size([5, 5, 1]),因为它为A 中的每个元素(或Af 中的)存储了B 中每个分位数的阈值.由于我们有 5 个分位数值,因此我们为 A 中的每个元素获得 5 个阈值。

    例如,quantiles[i, j, 0] 具有B[i]th 分位数A[j]Af[j] 的阈值,并且您基本上需要quantiles[k, k, 0] 的值在批处理大小范围内或此处为5。

    现在,为了满足您需要 B 中的对应分位数和 A 中的元素的阈值的要求,只需索引 quantiles 中的对角线元素并填充形状为 quantstorch.Size([5])

    最后要获得mask,将A 与每个元素的相应阈值进行比较。请注意,这使用与阈值进行广播的元素比较。 mask 具有torch.Size([5, 2, 2]) 所需的形状。

    【讨论】:

    • 虽然这段代码运行的很完美,但是可以进一步优化吗?
    • 我会调查的。但是您是否检查过它是否是您管道中的瓶颈?
    • 我只是 timeit 使用了我之前遇到的问题中的代码,而您的方法慢了大约 0.15 毫秒,我认为这不是我的代码运行缓慢的原因。又是 Tyvm!
    猜你喜欢
    • 2021-07-19
    • 2018-05-15
    • 2019-10-15
    • 2019-05-17
    • 1970-01-01
    • 2019-12-03
    • 2019-04-16
    • 1970-01-01
    • 1970-01-01
    相关资源
    最近更新 更多