mirror of https://github.com/kritiksoman/GIMP-ML
You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
21 lines
658 B
Python
21 lines
658 B
Python
### Copyright (C) 2017 NVIDIA Corporation. All rights reserved.
|
|
### Licensed under the CC BY-NC-SA 4.0 license (https://creativecommons.org/licenses/by-nc-sa/4.0/legalcode).
|
|
import torch
|
|
|
|
def create_model(opt):
|
|
if opt.model == 'pix2pixHD':
|
|
from .pix2pixHD_model import Pix2PixHDModel, InferenceModel
|
|
if opt.isTrain:
|
|
model = Pix2PixHDModel()
|
|
else:
|
|
model = InferenceModel()
|
|
|
|
model.initialize(opt)
|
|
if opt.verbose:
|
|
print("model [%s] was created" % (model.name()))
|
|
|
|
if opt.isTrain and len(opt.gpu_ids):
|
|
model = torch.nn.DataParallel(model, device_ids=opt.gpu_ids)
|
|
|
|
return model
|