(self)
| 330 | self.load_model() |
| 331 | |
| 332 | def load_model(self): |
| 333 | if self.config.load_weight != "": |
| 334 | trainable_states = torch.load(self.config.load_weight, map_location=torch.device("cpu")) |
| 335 | load_result = self.model.load_state_dict(trainable_states, strict=False) |
| 336 | assert ( |
| 337 | len(load_result.unexpected_keys) == 0 |
| 338 | ), f"Load model failed, unexpected keys {load_result.unexpected_keys.__str__()}" |
| 339 | |
| 340 | def save_model(self, finish=False): |
| 341 | if self.config.save_model and (finish or self._last_global_step_saved != self.global_step): |
no outgoing calls
no test coverage detected