MCPcopy Create free account
hub / github.com/QWTforGithub/T2LDM / build_model

Function build_model

eval/__init__.py:41–64  ·  view source on GitHub ↗
(dataset_name, model_name, device='cpu')

Source from the content-addressed store, hash-verified

39 'nuscenes': {'size': [32, 1024], 'fov': [3, -25], 'depth_range': [0.01, 50.0]}}
40
41def build_model(dataset_name, model_name, device='cpu'):
42 # config
43 model_folder = os.path.join(DEFAULT_ROOT, dataset_name, model_name)
44
45 if not os.path.isdir(model_folder):
46 raise Exception('Not Available Pretrained Weights!')
47
48 config = yaml.safe_load(open(os.path.join(model_folder, 'config.yaml'), 'r'))
49 if model_name != 'rangenet':
50 config = common.dict2namespace(config)
51
52 # build model
53 model = eval(model_name)(config)
54
55 # load checkpoint
56 if model_name == 'rangenet':
57 model.load_pretrained_weights(model_folder)
58 else:
59 ckpt = torch.load(os.path.join(model_folder, 'model.ckpt'), map_location="cpu")
60 model.load_state_dict(ckpt['state_dict'], strict=False)
61 model.to(device)
62 model.eval()
63
64 return model

Callers 1

compute_logitsFunction · 0.90

Calls 2

load_state_dictMethod · 0.45

Tested by

no test coverage detected