PyTorch Implementation of "Gradient Surgery for Multi-Task Learning" using multiprocessing
importtorchimporttorch.nnasnnimporttorch.optimasoptimfromppcgradimportPPCGrad# wrap your favorite optimizeroptimizer=PPCGrad(optim.Adam(net.parameters())) losses= [...] # a list of per-task lossesassertlen(losses) ==num_tasksoptimizer.pc_backward(losses) # calculate the gradient can apply gradient modificationoptimizer.step() # apply gradient step