[sgd] Add support for multi-model multi-optimizer training (#6317)

This commit is contained in:
Richard Liaw
2019-12-15 15:19:45 -08:00
committed by GitHub
parent c2499c802f
commit 5719a05757
12 changed files with 646 additions and 50 deletions
+2 -2
View File
@@ -77,7 +77,7 @@ def test(model, device, test_loader):
}
def dataset_creators(use_cuda):
def dataset_creator(use_cuda):
kwargs = {"num_workers": 1, "pin_memory": True} if use_cuda else {}
with FileLock("./data.lock"):
train_loader = torch.utils.data.DataLoader(
@@ -117,7 +117,7 @@ class Network(object):
def __init__(self, lr=0.01, momentum=0.5):
use_cuda = torch.cuda.is_available()
self.device = device = torch.device("cuda" if use_cuda else "cpu")
self.train_loader, self.test_loader = dataset_creators(use_cuda)
self.train_loader, self.test_loader = dataset_creator(use_cuda)
self.model = Model().to(device)
self.optimizer = optim.SGD(