MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / reset

Method reset

Image_Classification/src/train.py:202–246  ·  view source on GitHub ↗
(self, learning_rate: float)

Source from the content-addressed store, hash-verified

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):

Callers 1

trainMethod · 0.95

Calls 4

get_pretrained_modelMethod · 0.95
get_networkFunction · 0.90
get_optimizer_schedulerFunction · 0.90
resetMethod · 0.45

Tested by

no test coverage detected