PyTorch: Inference on a single very large image using multiple GPUs?

Viewed 422

I want to perform inference (i.e. semantic segmentation) on a very large satellite image without splitting it into pieces. I have access to 4 GPUs (each having 15 GBs of memory) and was wondering if it is possible to somehow use all the memory of these GPUs combined (i.e. 60 GB) for inference on the image in PyTorch?

1 Answers

You are looking for model parallel mode of work.
Basically, you can assign different parts of your model to different GPUs and then you should take care of the "bookkeeping".
This solution is very model-specific and task-specific therefore, there are no "generic" wrappers for it (as opposed to data parallel).

For example:

class MyModelParallelNetwork(nn.Module):
  def __init__(self, ...):
    # define the network
    self.part_one = ... # some nn.Module
    self.part_two = ... # additional nn.Module 
    self.part_three = ... 
    self.part_four = ...

    # important part - "send" the different parts to different GPUs
    self.part_one.to(torch.device('gpu:0'))
    self.part_two.to(torch.device('gpu:1'))
    self.part_three.to(torch.device('gpu:2'))
    self.part_four.to(torch.device('gpu:3'))

  def forward(self, x):
    # forward through model parts and GPUs:
    p1 = self.part_one(x.to(torch.device('gpu:0')))
    p2 = self.part_two(p1.to(torch.device('gpu:1')))
    p3 = self.part_three(p2.to(torch.device('gpu:2')))
    y = self.part_four(p3.to(torch.device('gpu:3')))
    return y  # result is on cuda:3 device
Related