Using PyTorch in multiple independent forked threads

Viewed 521

I need to run multiple threads in parallel, each solving a task using PyTorch. Importantly, those instances don't need to share any data computed by PyTorch, so I would expect it not to care about the parallelism at all. However, this seems to have changed between PyTorch 1.4 and 1.5.

Using PyTorch 1.4, the following minimal working example runs successfully:

import multiprocessing as mp
import torch


def train():
    print("Training started")
    x = torch.Tensor(0)
    x.requires_grad = True
    x.sum().backward()
    print("Training ended")


if __name__ == "__main__":
    print("Entered main")
    train()
    worker = mp.Process(target=train)
    worker.start()

And I get the output

Entered main
Training started
Training ended
Training started
Training ended

When I upgrade to PyTorch 1.5 or above, I get the following output instead:

Entered main
Training started
Training ended
Training started
Process Process-1:
Traceback (most recent call last):
  File "/usr/lib/python3.8/multiprocessing/process.py", line 315, in _bootstrap
    self.run()
  File "/usr/lib/python3.8/multiprocessing/process.py", line 108, in run
    self._target(*self._args, **self._kwargs)
  File "mwe.py", line 9, in train
    x.sum().backward()
  File "/home/christopher/.local/share/virtualenvs/src-i_X_I5Sj/lib/python3.8/site-packages/torch/tensor.py", line 198, in backward
    torch.autograd.backward(self, gradient, retain_graph, create_graph)
  File "/home/christopher/.local/share/virtualenvs/src-i_X_I5Sj/lib/python3.8/site-packages/torch/autograd/__init__.py", line 98, in backward
    Variable._execution_engine.run_backward(
RuntimeError: Unable to handle autograd's threading in combination with fork-based multiprocessing. See https://github.com/pytorch/pytorch/wiki/Autograd-and-Fork

As described by the given link, I can fix it by switching to worker = mp.get_context("spawn").Process(target=train). However, I would prefer not to do so, as fork is supposed to be faster than spawn.

Is it possible to encapsulate the calls to PyTorch in such a way that I can still use fork for multiprocessing?

0 Answers
Related