Skip to content

Commit eceb3a4

Browse files
committed
Fix loading CUDA-trained models on Mac MPS
1 parent a083f43 commit eceb3a4

1 file changed

Lines changed: 7 additions & 17 deletions

File tree

‎trainer/src/model_utils.py‎

Lines changed: 7 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -48,23 +48,13 @@ def get_latest_model_paths(model_dir, k):
4848

4949
def load_model(model_path):
5050
model = UNetGNRes()
51-
if torch.cuda.is_available() or torch.backends.mps.is_available():
52-
try:
53-
model.load_state_dict(torch.load(model_path))
54-
model = torch.nn.DataParallel(model)
55-
except:
56-
model = torch.nn.DataParallel(model)
57-
model.load_state_dict(torch.load(model_path))
58-
model.to(device)
59-
else:
60-
# if you are running on a CPU-only machine, please use torch.load with
61-
# map_location=torch.device('cpu') to map your storages to the CPU.
62-
try:
63-
model.load_state_dict(torch.load(model_path, map_location=torch.device('cpu')))
64-
model = torch.nn.DataParallel(model)
65-
except:
66-
model = torch.nn.DataParallel(model)
67-
model.load_state_dict(torch.load(model_path, map_location=torch.device('cpu')))
51+
try:
52+
model.load_state_dict(torch.load(model_path, map_location=device))
53+
model = torch.nn.DataParallel(model)
54+
except:
55+
model = torch.nn.DataParallel(model)
56+
model.load_state_dict(torch.load(model_path, map_location=device))
57+
model.to(device)
6858
return model
6959

7060
def create_first_model_with_random_weights(model_dir):

0 commit comments

Comments
 (0)