MCPcopy Create free account
hub / github.com/OpenDriveLab/DriveAdapter / load

Method load

roach/models/ppo_policy.py:222–240  ·  view source on GitHub ↗
(cls, path)

Source from the content-addressed store, hash-verified

220
221 @classmethod
222 def load(cls, path):
223 if th.cuda.is_available():
224 device = 'cuda'
225 else:
226 device = 'cpu'
227 saved_variables = th.load(path, map_location=device)
228 # Create policy object
229 model = cls(**saved_variables['policy_init_kwargs'])
230
231 state_dict = saved_variables['policy_state_dict']
232 for index in [2, 4, 6, 8, 10]:
233 state_dict["features_extractor.cnn."+str(index//2)+".weight"] = state_dict.pop("features_extractor.cnn."+str(index)+".weight")
234 state_dict["features_extractor.cnn."+str(index//2)+".bias"] = state_dict.pop("features_extractor.cnn."+str(index)+".bias")
235 #OLD!!! odict_keys(['features_extractor.cnn.0.weight', 'features_extractor.cnn.0.bias', 'features_extractor.cnn.2.weight', 'features_extractor.cnn.2.bias', 'features_extractor.cnn.4.weight', 'features_extractor.cnn.4.bias', 'features_extractor.cnn.6.weight', 'features_extractor.cnn.6.bias', 'features_extractor.cnn.8.weight', 'features_extractor.cnn.8.bias', 'features_extractor.cnn.10.weight', 'features_extractor.cnn.10.bias', 'features_extractor.linear.0.weight', 'features_extractor.linear.0.bias', 'features_extractor.linear.2.weight', 'features_extractor.linear.2.bias', 'features_extractor.state_linear.0.weight', 'features_extractor.state_linear.0.bias', 'features_extractor.state_linear.2.weight', 'features_extractor.state_linear.2.bias', 'policy_head.0.weight', 'policy_head.0.bias', 'policy_head.2.weight', 'policy_head.2.bias', 'dist_mu.0.weight', 'dist_mu.0.bias', 'dist_sigma.0.weight', 'dist_sigma.0.bias', 'value_head.0.weight', 'value_head.0.bias', 'value_head.2.weight', 'value_head.2.bias', 'value_head.4.weight', 'value_head.4.bias'])
236 # odict_keys(['features_extractor.cnn.0.weight', 'features_extractor.cnn.0.bias', 'features_extractor.cnn.1.weight', 'features_extractor.cnn.1.bias', 'features_extractor.cnn.2.weight', 'features_extractor.cnn.2.bias', 'features_extractor.cnn.3.weight', 'features_extractor.cnn.3.bias', 'features_extractor.cnn.4.weight', 'features_extractor.cnn.4.bias', 'features_extractor.cnn.5.weight', 'features_extractor.cnn.5.bias', 'features_extractor.linear.0.weight', 'features_extractor.linear.0.bias', 'features_extractor.linear.2.weight', 'features_extractor.linear.2.bias', 'features_extractor.state_linear.0.weight', 'features_extractor.state_linear.0.bias', 'features_extractor.state_linear.2.weight', 'features_extractor.state_linear.2.bias', 'policy_head.0.weight', 'policy_head.0.bias', 'policy_head.2.weight', 'policy_head.2.bias', 'dist_mu.0.weight', 'dist_mu.0.bias', 'dist_sigma.0.weight', 'dist_sigma.0.bias', 'value_head.0.weight', 'value_head.0.bias', 'value_head.2.weight', 'value_head.2.bias', 'value_head.4.weight', 'value_head.4.bias'])
237 # Load weights
238 model.load_state_dict(saved_variables['policy_state_dict'])
239 model.to(device)
240 return model, saved_variables['train_init_kwargs']
241
242 @staticmethod
243 def init_weights(module: nn.Module, gain: float = 1) -> None:

Callers 15

visualize_roach_inputMethod · 0.80
load_roach_layersMethod · 0.80
collect_results_cpuFunction · 0.80
load_annotationsMethod · 0.80
__init__Method · 0.80
load_jsonMethod · 0.80
load_npyMethod · 0.80
load_bev_segMethod · 0.80
_load_pointsMethod · 0.80
__init__Method · 0.80
fetch_dictFunction · 0.80

Calls

no outgoing calls

Tested by

no test coverage detected