(self, learning_rate: float)
| 200 | raise NotImplementedError(f"{self.config['network']} does not support pretrained weights yet") |
| 201 | |
| 202 | def reset(self, learning_rate: float) -> None: |
| 203 | self.network = get_network(name=self.config['network'], num_classes=self.num_classes) |
| 204 | for layer in self.network.modules(): |
| 205 | if isinstance(layer, nn.BatchNorm2d): |
| 206 | layer.float() |
| 207 | if self.pretrained: |
| 208 | self.get_pretrained_model() |
| 209 | if self.device == 'cpu': |
| 210 | print("Resetting cpu-based network") |
| 211 | elif self.device == 'mps': |
| 212 | self.network = self.network.to(self.gpu) |
| 213 | elif self.dist: |
| 214 | self.network.to(self.gpu) |
| 215 | self.network = torch.nn.parallel.DistributedDataParallel( |
| 216 | self.network, device_ids=[self.gpu]) |
| 217 | else: |
| 218 | self.network = self.network.cuda(self.gpu) |
| 219 | |
| 220 | self.optimizer, self.scheduler = get_optimizer_scheduler( |
| 221 | optim_method=self.config['optimizer'], |
| 222 | lr_scheduler=self.config['scheduler'], |
| 223 | init_lr=learning_rate, |
| 224 | net=self.network, |
| 225 | train_loader_len=len(self.train_loader), |
| 226 | max_epochs=self.config['max_epochs'], |
| 227 | optimizer_kwargs=self.config['optimizer_kwargs'], |
| 228 | scheduler_kwargs=self.config['scheduler_kwargs']) |
| 229 | if self.config['optimizer'] in ['Shampoo', 'kfac']: |
| 230 | damping = self.config['optimizer_kwargs']['damping'] |
| 231 | curvature_update_interval = self.config['optimizer_kwargs']['curvature_update_interval'] |
| 232 | ema_decay = self.config['optimizer_kwargs']['ema_decay'] |
| 233 | config = PreconditioningConfig(data_size=self.config['mini_batch_size'], |
| 234 | damping=damping, |
| 235 | preconditioner_upd_interval=curvature_update_interval, |
| 236 | curvature_upd_interval=curvature_update_interval, |
| 237 | ema_decay=ema_decay, |
| 238 | ignore_modules=[nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d, |
| 239 | nn.LayerNorm]) |
| 240 | if self.config['optimizer'] == "Shampoo": |
| 241 | self.gm = ShampooGradientMaker(self.network, config) |
| 242 | elif self.config['optimizer'] == "kfac": |
| 243 | self.gm = KfacGradientMaker(self.network, config) |
| 244 | else: |
| 245 | raise ValueError(f"This optimizer is not reconized: {self.config['optimizer']}") |
| 246 | self.early_stop.reset() |
| 247 | |
| 248 | def create_output_dir(self, checkpoint: bool = False): |
| 249 | def create_dir(path: Path): |
no test coverage detected