I want to sample a certain number of points in each region of the image according to a segment map(like what the SLIC algorithm produces, it's a map with the same size as the image, containing integers from 0 to num_segment indicating which segment each pixel belongs to).
Currently, I write my own code as follows:
- For
iin range(0, num_segment):- find the (indexes of the)pixels that belong to the
ith segment usingtorch.where - picking out those pixels and forming a 1-d tensor
- use
torch.Upsampleto uniformly sample n_sample points for theith segment
- find the (indexes of the)pixels that belong to the
- stack all the sampled points to form a large 2-d tensor which each row represent selected points belong to one segment, and it has n_sample rows.
I not only want the original value from the image for each selected point, but their indexes are also needed.
I drew a picture to illustrate.
So, my question is, is there a native way to implement this in PyTorch? The above code runs a little bit slow, maybe the For-loop slows down the speed. And if possible, since all the sampling process are independent of each other, how can I speed up the process?
Generally, I have ~400 segments, and I want to sample 20~50 points for each segment.