【问题标题】:How to combine/stack tensors and combine dimensions in PyTorch?如何在 PyTorch 中组合/堆叠张量和组合维度?
【发布时间】:2019-07-05 16:28:15
【问题描述】:

我需要将 4 个张量(代表灰度图像,大小为 [1,84,84])组合成一堆形状 [4,84,84],代表四个灰度图像,每个图像表示为“通道”张量样式 CxWxH。

我正在使用 PyTorch。

我已经尝试使用 torch.stack 和 torch.cat,但如果其中一个是解决方案,我没有运气找出正确的准备/方法来获得我的结果。

感谢您的帮助。

import torchvision.transforms as T

class ReplayBuffer:
    def __init__(self, buffersize, batchsize, framestack, device, nS):
        self.buffer = deque(maxlen=buffersize)
        self.phi = deque(maxlen=framestack)
        self.batchsize = batchsize
        self.device = device

        self._initialize_stack(nS)

    def get_stack(self):
        #t =  torch.cat(tuple(self.phi),dim=0)
        t =  torch.stack(tuple(self.phi),dim=0)
        return t

    def _initialize_stack(self, nS):
        while len(self.phi) < self.phi.maxlen:
            self.phi.append(torch.tensor([1,nS[1], nS[2]]))

a = ReplayBuffer(buffersize=50000, batchsize=64, framestack=4, device='cuda', nS=[1,84,84])
print(a.phi)
s = a.get_stack()
print(s, s.shape)

以上代码返回:

print(a.phi)

deque([tensor([ 1, 84, 84]), tensor([ 1, 84, 84]), tensor([ 1, 84, 84]), tensor([ 1, 84, 84])], maxlen=4)

print(s, s.shape)

tensor([[ 1, 84, 84],
        [ 1, 84, 84],
        [ 1, 84, 84],
        [ 1, 84, 84]]) torch.Size([4, 3])

但我想要的只是返回 [4, 84, 84]。我怀疑这很简单,但它让我无法理解。

【问题讨论】:

    标签: python stack concatenation pytorch tensor


    【解决方案1】:

    看来你误解了torch.tensor([1, 84, 84]) 的作用。一起来看看吧:

    torch.tensor([1, 84, 84])
    print(x, x.shape) #tensor([ 1, 84, 84]) torch.Size([3])
    

    从上面的例子可以看出,它给了你一个只有一维的张量。

    根据您的问题陈述,您需要一个形状为 [1,84,84] 的张量。 下面是它的样子:

    from collections import deque
    import torch
    import torchvision.transforms as T
    
    class ReplayBuffer:
        def __init__(self, buffersize, batchsize, framestack, device, nS):
            self.buffer = deque(maxlen=buffersize)
            self.phi = deque(maxlen=framestack)
            self.batchsize = batchsize
            self.device = device
    
            self._initialize_stack(nS)
    
        def get_stack(self):
            t =  torch.cat(tuple(self.phi),dim=0)
    #         t =  torch.stack(tuple(self.phi),dim=0)
            return t
    
        def _initialize_stack(self, nS):
            while len(self.phi) < self.phi.maxlen:
    #             self.phi.append(torch.tensor([1,nS[1], nS[2]]))
                self.phi.append(torch.zeros([1,nS[1], nS[2]]))
    
    a = ReplayBuffer(buffersize=50000, batchsize=64, framestack=4, device='cuda', nS=[1,84,84])
    print(a.phi)
    s = a.get_stack()
    print(s, s.shape)
    

    请注意,torch.cat 给你一个形状为 [4, 84, 84] 的张量,torch.stack 给你一个形状为 [4, 1, 84, 84] 的张量。他们的区别可以在What's the difference between torch.stack() and torch.cat() functions?找到。

    【讨论】:

    • 非常感谢!我知道这将是一件相当简单的事情,但我不能说我曾经设法以我目前的知识来确定那个错误。我将继续研究张量创建/处理。干杯。
    • @White_Rabbit.obj 如果这个答案对你有帮助,请考虑采纳,谢谢!
    猜你喜欢
    • 2019-12-31
    • 2020-12-23
    • 1970-01-01
    • 2021-01-10
    • 1970-01-01
    • 1970-01-01
    • 2021-04-06
    • 1970-01-01
    • 2021-10-23
    相关资源
    最近更新 更多