Pytorch broadcasting command not found

Viewed 68

I have the following segment of nested for loop in my code. The nested loop is slowing down my complete execution.

for a torch tensor extended_output with shape [batchSize,nClass*repeat] and another torch tensor with dimension [batchSize,nClass], I want the aggregation to happen as follows:

for q in range(nClass):
    for u in range(repeat):
        output[:,q]=output[:,q]+extended_output[:,(q+u*nClass)]

Here, nClass,repeat all are integer variables with value 1400 and 8 repectively.

Can this nested for loop be avoided using pytorch broadcasting? Any help will be highly useful.

A sample working cpode might be like this

import torch
nClass=1400
repeat=8
batchSize=64
output=torch.zeros([batchSize,nClass])
extended_output=torch.rand([batchSize,nClass*repeat])

for q in range(nClass):
    for u in range(repeat):
        output[:,q]=output[:,q]+extended_output[:,(q+u*nClass)]
1 Answers

Sorry for the short and probably over-simplified example. I fear a bigger one would be much more difficult to visualize. But I hope this suits your purpose. Here's what I would do:

import torch
nClass    = 3
repeat    = 2
batchSize = 4

torch.manual_seed(0)

output          = torch.zeros([batchSize,nClass])
extended_output = torch.rand([batchSize,nClass*repeat])


for q in range(nClass):
    for u in range(repeat):
        output[:,q]=output[:,q]+extended_output[:,(q+u*nClass)]

idxs = (torch.arange(repeat)*nClass).unsqueeze(0)
idxs = idxs + torch.arange(nClass).unsqueeze(1)
output_vectorized = extended_output[:, idxs].sum(2)

output:

extended_output = 
tensor([[0.4963, 0.7682, 0.0885, 0.1320, 0.3074, 0.6341],
        [0.4901, 0.8964, 0.4556, 0.6323, 0.3489, 0.4017],
        [0.0223, 0.1689, 0.2939, 0.5185, 0.6977, 0.8000],
        [0.1610, 0.2823, 0.6816, 0.9152, 0.3971, 0.8742]])
output = 
tensor([[0.6283, 1.0756, 0.7226],
        [1.1224, 1.2453, 0.8573],
        [0.5408, 0.8665, 1.0939],
        [1.0762, 0.6794, 1.5558]])
output_vectorized = 
tensor([[0.6283, 1.0756, 0.7226],
        [1.1224, 1.2453, 0.8573],
        [0.5408, 0.8665, 1.0939],
        [1.0762, 0.6794, 1.5558]])
Related