MCPcopy Create free account
hub / github.com/OpenDriveLab/TCP / _load_state_dict

Method _load_state_dict

TCP/train.py:36–43  ·  view source on GitHub ↗
(self, il_net, rl_state_dict, key_word)

Source from the content-addressed store, hash-verified

34 self._load_state_dict(self.model.dist_sigma, rl_state_dict, 'dist_sigma')
35
36 def _load_state_dict(self, il_net, rl_state_dict, key_word):
37 rl_keys = [k for k in rl_state_dict.keys() if key_word in k]
38 il_keys = il_net.state_dict().keys()
39 assert len(rl_keys) == len(il_net.state_dict().keys()), f'mismatch number of layers loading {key_word}'
40 new_state_dict = OrderedDict()
41 for k_il, k_rl in zip(il_keys, rl_keys):
42 new_state_dict[k_il] = rl_state_dict[k_rl]
43 il_net.load_state_dict(new_state_dict)
44
45 def forward(self, batch):
46 pass

Callers 1

_load_weightMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected