Arulkumar03/Wheat_HEAD_Detection_Counting_ComputerVision_Model
0
1# -*- coding: utf-8 -*-2# Copyright (c) Facebook, Inc. and its affiliates.3 4"""5This file contains components with some default boilerplate logic user may need6in training / testing. They will not work for everyone, but many users may find them useful.7 8The behavior of functions/classes in this file is subject to change,9since they are meant to represent the "common default behavior" people need in their projects.10"""11 12import argparse13import logging14import os15import sys16import weakref17from collections import OrderedDict18from typing import Optional19import torch20from fvcore.nn.precise_bn import get_bn_modules21from omegaconf import OmegaConf22from torch.nn.parallel import DistributedDataParallel23 24import detectron2.data.transforms as T25from detectron2.checkpoint import DetectionCheckpointer26from detectron2.config import CfgNode, LazyConfig27from detectron2.data import (28 MetadataCatalog,29 build_detection_test_loader,30 build_detection_train_loader,31)32from detectron2.evaluation import (33 DatasetEvaluator,34 inference_on_dataset,35 print_csv_format,36 verify_results,37)38from detectron2.modeling import build_model39from detectron2.solver import build_lr_scheduler, build_optimizer40from detectron2.utils import comm41from detectron2.utils.collect_env import collect_env_info42from detectron2.utils.env import seed_all_rng43from detectron2.utils.events import CommonMetricPrinter, JSONWriter, TensorboardXWriter44from detectron2.utils.file_io import PathManager45from detectron2.utils.logger import setup_logger46 47from . import hooks48from .train_loop import AMPTrainer, SimpleTrainer, TrainerBase49 50__all__ = [51 "create_ddp_model",52 "default_argument_parser",53 "default_setup",54 "default_writers",55 "DefaultPredictor",56 "DefaultTrainer",57]58 59 60def create_ddp_model(model, *, fp16_compression=False, **kwargs):61 """62 Create a DistributedDataParallel model if there are >1 processes.63 64 Args:65 model: a torch.nn.Module66 fp16_compression: add fp16 compression hooks to the ddp object.67 See more at https://pytorch.org/docs/stable/ddp_comm_hooks.html#torch.distributed.algorithms.ddp_comm_hooks.default_hooks.fp16_compress_hook68 kwargs: other arguments of :module:`torch.nn.parallel.DistributedDataParallel`.69 """ # noqa70 if comm.get_world_size() == 1:71 return model72 if "device_ids" not in kwargs:73 kwargs["device_ids"] = [comm.get_local_rank()]74 ddp = DistributedDataParallel(model, **kwargs)75 if fp16_compression:76 from torch.distributed.algorithms.ddp_comm_hooks import default as comm_hooks77 78 ddp.register_comm_hook(state=None, hook=comm_hooks.fp16_compress_hook)79 return ddp80 81 82def default_argument_parser(epilog=None):83 """84 Create a parser with some common arguments used by detectron2 users.85 86 Args:87 epilog (str): epilog passed to ArgumentParser describing the usage.88 89 Returns:90 argparse.ArgumentParser:91 """92 parser = argparse.ArgumentParser(93 epilog=epilog94 or f"""95Examples:96 97Run on single machine:98 $ {sys.argv[0]} --num-gpus 8 --config-file cfg.yaml99 100Change some config options:101 $ {sys.argv[0]} --config-file cfg.yaml MODEL.WEIGHTS /path/to/weight.pth SOLVER.BASE_LR 0.001102 103Run on multiple machines:104 (machine0)$ {sys.argv[0]} --machine-rank 0 --num-machines 2 --dist-url <URL> [--other-flags]105 (machine1)$ {sys.argv[0]} --machine-rank 1 --num-machines 2 --dist-url <URL> [--other-flags]106""",107 formatter_class=argparse.RawDescriptionHelpFormatter,108 )109 parser.add_argument("--config-file", default="", metavar="FILE", help="path to config file")110 parser.add_argument(111 "--resume",112 action="store_true",113 help="Whether to attempt to resume from the checkpoint directory. "114 "See documentation of `DefaultTrainer.resume_or_load()` for what it means.",115 )116 parser.add_argument("--eval-only", action="store_true", help="perform evaluation only")117 parser.add_argument("--num-gpus", type=int, default=1, help="number of gpus *per machine*")118 parser.add_argument("--num-machines", type=int, default=1, help="total number of machines")119 parser.add_argument(120 "--machine-rank", type=int, default=0, help="the rank of this machine (unique per machine)"121 )122 123 # PyTorch still may leave orphan processes in multi-gpu training.124 # Therefore we use a deterministic way to obtain port,125 # so that users are aware of orphan processes by seeing the port occupied.126 port = 2**15 + 2**14 + hash(os.getuid() if sys.platform != "win32" else 1) % 2**14127 parser.add_argument(128 "--dist-url",129 default="tcp://127.0.0.1:{}".format(port),130 help="initialization URL for pytorch distributed backend. See "131 "https://pytorch.org/docs/stable/distributed.html for details.",132 )133 parser.add_argument(134 "opts",135 help="""136Modify config options at the end of the command. For Yacs configs, use137space-separated "PATH.KEY VALUE" pairs.138For python-based LazyConfig, use "path.key=value".139 """.strip(),140 default=None,141 nargs=argparse.REMAINDER,142 )143 return parser144 145 146def _try_get_key(cfg, *keys, default=None):147 """148 Try select keys from cfg until the first key that exists. Otherwise return default.149 """150 if isinstance(cfg, CfgNode):151 cfg = OmegaConf.create(cfg.dump())152 for k in keys:153 none = object()154 p = OmegaConf.select(cfg, k, default=none)155 if p is not none:156 return p157 return default158 159 160def _highlight(code, filename):161 try:162 import pygments163 except ImportError:164 return code165 166 from pygments.lexers import Python3Lexer, YamlLexer167 from pygments.formatters import Terminal256Formatter168 169 lexer = Python3Lexer() if filename.endswith(".py") else YamlLexer()170 code = pygments.highlight(code, lexer, Terminal256Formatter(style="monokai"))171 return code172 173 174def default_setup(cfg, args):175 """176 Perform some basic common setups at the beginning of a job, including:177 178 1. Set up the detectron2 logger179 2. Log basic information about environment, cmdline arguments, and config180 3. Backup the config to the output directory181 182 Args:183 cfg (CfgNode or omegaconf.DictConfig): the full config to be used184 args (argparse.NameSpace): the command line arguments to be logged185 """186 output_dir = _try_get_key(cfg, "OUTPUT_DIR", "output_dir", "train.output_dir")187 if comm.is_main_process() and output_dir:188 PathManager.mkdirs(output_dir)189 190 rank = comm.get_rank()191 setup_logger(output_dir, distributed_rank=rank, name="fvcore")192 logger = setup_logger(output_dir, distributed_rank=rank)193 194 logger.info("Rank of current process: {}. World size: {}".format(rank, comm.get_world_size()))195 logger.info("Environment info:\n" + collect_env_info())196 197 logger.info("Command line arguments: " + str(args))198 if hasattr(args, "config_file") and args.config_file != "":199 logger.info(200 "Contents of args.config_file={}:\n{}".format(201 args.config_file,202 _highlight(PathManager.open(args.config_file, "r").read(), args.config_file),203 )204 )205 206 if comm.is_main_process() and output_dir:207 # Note: some of our scripts may expect the existence of208 # config.yaml in output directory209 path = os.path.join(output_dir, "config.yaml")210 if isinstance(cfg, CfgNode):211 logger.info("Running with full config:\n{}".format(_highlight(cfg.dump(), ".yaml")))212 with PathManager.open(path, "w") as f:213 f.write(cfg.dump())214 else:215 LazyConfig.save(cfg, path)216 logger.info("Full config saved to {}".format(path))217 218 # make sure each worker has a different, yet deterministic seed if specified219 seed = _try_get_key(cfg, "SEED", "train.seed", default=-1)220 seed_all_rng(None if seed < 0 else seed + rank)221 222 # cudnn benchmark has large overhead. It shouldn't be used considering the small size of223 # typical validation set.224 if not (hasattr(args, "eval_only") and args.eval_only):225 torch.backends.cudnn.benchmark = _try_get_key(226 cfg, "CUDNN_BENCHMARK", "train.cudnn_benchmark", default=False227 )228 229 230def default_writers(output_dir: str, max_iter: Optional[int] = None):231 """232 Build a list of :class:`EventWriter` to be used.233 It now consists of a :class:`CommonMetricPrinter`,234 :class:`TensorboardXWriter` and :class:`JSONWriter`.235 236 Args:237 output_dir: directory to store JSON metrics and tensorboard events238 max_iter: the total number of iterations239 240 Returns:241 list[EventWriter]: a list of :class:`EventWriter` objects.242 """243 PathManager.mkdirs(output_dir)244 return [245 # It may not always print what you want to see, since it prints "common" metrics only.246 CommonMetricPrinter(max_iter),247 JSONWriter(os.path.join(output_dir, "metrics.json")),248 TensorboardXWriter(output_dir),249 ]250 251 252class DefaultPredictor:253 """254 Create a simple end-to-end predictor with the given config that runs on255 single device for a single input image.256 257 Compared to using the model directly, this class does the following additions:258 259 1. Load checkpoint from `cfg.MODEL.WEIGHTS`.260 2. Always take BGR image as the input and apply conversion defined by `cfg.INPUT.FORMAT`.261 3. Apply resizing defined by `cfg.INPUT.{MIN,MAX}_SIZE_TEST`.262 4. Take one input image and produce a single output, instead of a batch.263 264 This is meant for simple demo purposes, so it does the above steps automatically.265 This is not meant for benchmarks or running complicated inference logic.266 If you'd like to do anything more complicated, please refer to its source code as267 examples to build and use the model manually.268 269 Attributes:270 metadata (Metadata): the metadata of the underlying dataset, obtained from271 cfg.DATASETS.TEST.272 273 Examples:274 ::275 pred = DefaultPredictor(cfg)276 inputs = cv2.imread("input.jpg")277 outputs = pred(inputs)278 """279 280 def __init__(self, cfg):281 self.cfg = cfg.clone() # cfg can be modified by model282 self.model = build_model(self.cfg)283 self.model.eval()284 if len(cfg.DATASETS.TEST):285 self.metadata = MetadataCatalog.get(cfg.DATASETS.TEST[0])286 287 checkpointer = DetectionCheckpointer(self.model)288 checkpointer.load(cfg.MODEL.WEIGHTS)289 290 self.aug = T.ResizeShortestEdge(291 [cfg.INPUT.MIN_SIZE_TEST, cfg.INPUT.MIN_SIZE_TEST], cfg.INPUT.MAX_SIZE_TEST292 )293 294 self.input_format = cfg.INPUT.FORMAT295 assert self.input_format in ["RGB", "BGR"], self.input_format296 297 def __call__(self, original_image):298 """299 Args:300 original_image (np.ndarray): an image of shape (H, W, C) (in BGR order).301 302 Returns:303 predictions (dict):304 the output of the model for one image only.305 See :doc:`/tutorials/models` for details about the format.306 """307 with torch.no_grad(): # https://github.com/sphinx-doc/sphinx/issues/4258308 # Apply pre-processing to image.309 if self.input_format == "RGB":310 # whether the model expects BGR inputs or RGB311 original_image = original_image[:, :, ::-1]312 height, width = original_image.shape[:2]313 image = self.aug.get_transform(original_image).apply_image(original_image)314 image = torch.as_tensor(image.astype("float32").transpose(2, 0, 1))315 316 inputs = {"image": image, "height": height, "width": width}317 predictions = self.model([inputs])[0]318 return predictions319 320 321class DefaultTrainer(TrainerBase):322 """323 A trainer with default training logic. It does the following:324 325 1. Create a :class:`SimpleTrainer` using model, optimizer, dataloader326 defined by the given config. Create a LR scheduler defined by the config.327 2. Load the last checkpoint or `cfg.MODEL.WEIGHTS`, if exists, when328 `resume_or_load` is called.329 3. Register a few common hooks defined by the config.330 331 It is created to simplify the **standard model training workflow** and reduce code boilerplate332 for users who only need the standard training workflow, with standard features.333 It means this class makes *many assumptions* about your training logic that334 may easily become invalid in a new research. In fact, any assumptions beyond those made in the335 :class:`SimpleTrainer` are too much for research.336 337 The code of this class has been annotated about restrictive assumptions it makes.338 When they do not work for you, you're encouraged to:339 340 1. Overwrite methods of this class, OR:341 2. Use :class:`SimpleTrainer`, which only does minimal SGD training and342 nothing else. You can then add your own hooks if needed. OR:343 3. Write your own training loop similar to `tools/plain_train_net.py`.344 345 See the :doc:`/tutorials/training` tutorials for more details.346 347 Note that the behavior of this class, like other functions/classes in348 this file, is not stable, since it is meant to represent the "common default behavior".349 It is only guaranteed to work well with the standard models and training workflow in detectron2.350 To obtain more stable behavior, write your own training logic with other public APIs.351 352 Examples:353 ::354 trainer = DefaultTrainer(cfg)355 trainer.resume_or_load() # load last checkpoint or MODEL.WEIGHTS356 trainer.train()357 358 Attributes:359 scheduler:360 checkpointer (DetectionCheckpointer):361 cfg (CfgNode):362 """363 364 def __init__(self, cfg):365 """366 Args:367 cfg (CfgNode):368 """369 super().__init__()370 logger = logging.getLogger("detectron2")371 if not logger.isEnabledFor(logging.INFO): # setup_logger is not called for d2372 setup_logger()373 cfg = DefaultTrainer.auto_scale_workers(cfg, comm.get_world_size())374 375 # Assume these objects must be constructed in this order.376 model = self.build_model(cfg)377 optimizer = self.build_optimizer(cfg, model)378 data_loader = self.build_train_loader(cfg)379 380 model = create_ddp_model(model, broadcast_buffers=False)381 self._trainer = (AMPTrainer if cfg.SOLVER.AMP.ENABLED else SimpleTrainer)(382 model, data_loader, optimizer383 )384 385 self.scheduler = self.build_lr_scheduler(cfg, optimizer)386 self.checkpointer = DetectionCheckpointer(387 # Assume you want to save checkpoints together with logs/statistics388 model,389 cfg.OUTPUT_DIR,390 trainer=weakref.proxy(self),391 )392 self.start_iter = 0393 self.max_iter = cfg.SOLVER.MAX_ITER394 self.cfg = cfg395 396 self.register_hooks(self.build_hooks())397 398 def resume_or_load(self, resume=True):399 """400 If `resume==True` and `cfg.OUTPUT_DIR` contains the last checkpoint (defined by401 a `last_checkpoint` file), resume from the file. Resuming means loading all402 available states (eg. optimizer and scheduler) and update iteration counter403 from the checkpoint. ``cfg.MODEL.WEIGHTS`` will not be used.404 405 Otherwise, this is considered as an independent training. The method will load model406 weights from the file `cfg.MODEL.WEIGHTS` (but will not load other states) and start407 from iteration 0.408 409 Args:410 resume (bool): whether to do resume or not411 """412 self.checkpointer.resume_or_load(self.cfg.MODEL.WEIGHTS, resume=resume)413 if resume and self.checkpointer.has_checkpoint():414 # The checkpoint stores the training iteration that just finished, thus we start415 # at the next iteration416 self.start_iter = self.iter + 1417 418 def build_hooks(self):419 """420 Build a list of default hooks, including timing, evaluation,421 checkpointing, lr scheduling, precise BN, writing events.422 423 Returns:424 list[HookBase]:425 """426 cfg = self.cfg.clone()427 cfg.defrost()428 cfg.DATALOADER.NUM_WORKERS = 0 # save some memory and time for PreciseBN429 430 ret = [431 hooks.IterationTimer(),432 hooks.LRScheduler(),433 hooks.PreciseBN(434 # Run at the same freq as (but before) evaluation.435 cfg.TEST.EVAL_PERIOD,436 self.model,437 # Build a new data loader to not affect training438 self.build_train_loader(cfg),439 cfg.TEST.PRECISE_BN.NUM_ITER,440 )441 if cfg.TEST.PRECISE_BN.ENABLED and get_bn_modules(self.model)442 else None,443 ]444 445 # Do PreciseBN before checkpointer, because it updates the model and need to446 # be saved by checkpointer.447 # This is not always the best: if checkpointing has a different frequency,448 # some checkpoints may have more precise statistics than others.449 if comm.is_main_process():450 ret.append(hooks.PeriodicCheckpointer(self.checkpointer, cfg.SOLVER.CHECKPOINT_PERIOD))451 452 def test_and_save_results():453 self._last_eval_results = self.test(self.cfg, self.model)454 return self._last_eval_results455 456 # Do evaluation after checkpointer, because then if it fails,457 # we can use the saved checkpoint to debug.458 ret.append(hooks.EvalHook(cfg.TEST.EVAL_PERIOD, test_and_save_results))459 460 if comm.is_main_process():461 # Here the default print/log frequency of each writer is used.462 # run writers in the end, so that evaluation metrics are written463 ret.append(hooks.PeriodicWriter(self.build_writers(), period=20))464 return ret465 466 def build_writers(self):467 """468 Build a list of writers to be used using :func:`default_writers()`.469 If you'd like a different list of writers, you can overwrite it in470 your trainer.471 472 Returns:473 list[EventWriter]: a list of :class:`EventWriter` objects.474 """475 return default_writers(self.cfg.OUTPUT_DIR, self.max_iter)476 477 def train(self):478 """479 Run training.480 481 Returns:482 OrderedDict of results, if evaluation is enabled. Otherwise None.483 """484 super().train(self.start_iter, self.max_iter)485 if len(self.cfg.TEST.EXPECTED_RESULTS) and comm.is_main_process():486 assert hasattr(487 self, "_last_eval_results"488 ), "No evaluation results obtained during training!"489 verify_results(self.cfg, self._last_eval_results)490 return self._last_eval_results491 492 def run_step(self):493 self._trainer.iter = self.iter494 self._trainer.run_step()495 496 def state_dict(self):497 ret = super().state_dict()498 ret["_trainer"] = self._trainer.state_dict()499 return ret500 501 def load_state_dict(self, state_dict):502 super().load_state_dict(state_dict)503 self._trainer.load_state_dict(state_dict["_trainer"])504 505 @classmethod506 def build_model(cls, cfg):507 """508 Returns:509 torch.nn.Module:510 511 It now calls :func:`detectron2.modeling.build_model`.512 Overwrite it if you'd like a different model.513 """514 model = build_model(cfg)515 logger = logging.getLogger(__name__)516 logger.info("Model:\n{}".format(model))517 return model518 519 @classmethod520 def build_optimizer(cls, cfg, model):521 """522 Returns:523 torch.optim.Optimizer:524 525 It now calls :func:`detectron2.solver.build_optimizer`.526 Overwrite it if you'd like a different optimizer.527 """528 return build_optimizer(cfg, model)529 530 @classmethod531 def build_lr_scheduler(cls, cfg, optimizer):532 """533 It now calls :func:`detectron2.solver.build_lr_scheduler`.534 Overwrite it if you'd like a different scheduler.535 """536 return build_lr_scheduler(cfg, optimizer)537 538 @classmethod539 def build_train_loader(cls, cfg):540 """541 Returns:542 iterable543 544 It now calls :func:`detectron2.data.build_detection_train_loader`.545 Overwrite it if you'd like a different data loader.546 """547 return build_detection_train_loader(cfg)548 549 @classmethod550 def build_test_loader(cls, cfg, dataset_name):551 """552 Returns:553 iterable554 555 It now calls :func:`detectron2.data.build_detection_test_loader`.556 Overwrite it if you'd like a different data loader.557 """558 return build_detection_test_loader(cfg, dataset_name)559 560 @classmethod561 def build_evaluator(cls, cfg, dataset_name):562 """563 Returns:564 DatasetEvaluator or None565 566 It is not implemented by default.567 """568 raise NotImplementedError(569 """570If you want DefaultTrainer to automatically run evaluation,571please implement `build_evaluator()` in subclasses (see train_net.py for example).572Alternatively, you can call evaluation functions yourself (see Colab balloon tutorial for example).573"""574 )575 576 @classmethod577 def test(cls, cfg, model, evaluators=None):578 """579 Evaluate the given model. The given model is expected to already contain580 weights to evaluate.581 582 Args:583 cfg (CfgNode):584 model (nn.Module):585 evaluators (list[DatasetEvaluator] or None): if None, will call586 :meth:`build_evaluator`. Otherwise, must have the same length as587 ``cfg.DATASETS.TEST``.588 589 Returns:590 dict: a dict of result metrics591 """592 logger = logging.getLogger(__name__)593 if isinstance(evaluators, DatasetEvaluator):594 evaluators = [evaluators]595 if evaluators is not None:596 assert len(cfg.DATASETS.TEST) == len(evaluators), "{} != {}".format(597 len(cfg.DATASETS.TEST), len(evaluators)598 )599 600 results = OrderedDict()601 for idx, dataset_name in enumerate(cfg.DATASETS.TEST):602 data_loader = cls.build_test_loader(cfg, dataset_name)603 # When evaluators are passed in as arguments,604 # implicitly assume that evaluators can be created before data_loader.605 if evaluators is not None:606 evaluator = evaluators[idx]607 else:608 try:609 evaluator = cls.build_evaluator(cfg, dataset_name)610 except NotImplementedError:611 logger.warn(612 "No evaluator found. Use `DefaultTrainer.test(evaluators=)`, "613 "or implement its `build_evaluator` method."614 )615 results[dataset_name] = {}616 continue617 results_i = inference_on_dataset(model, data_loader, evaluator)618 results[dataset_name] = results_i619 if comm.is_main_process():620 assert isinstance(621 results_i, dict622 ), "Evaluator must return a dict on the main process. Got {} instead.".format(623 results_i624 )625 logger.info("Evaluation results for {} in csv format:".format(dataset_name))626 print_csv_format(results_i)627 628 if len(results) == 1:629 results = list(results.values())[0]630 return results631 632 @staticmethod633 def auto_scale_workers(cfg, num_workers: int):634 """635 When the config is defined for certain number of workers (according to636 ``cfg.SOLVER.REFERENCE_WORLD_SIZE``) that's different from the number of637 workers currently in use, returns a new cfg where the total batch size638 is scaled so that the per-GPU batch size stays the same as the639 original ``IMS_PER_BATCH // REFERENCE_WORLD_SIZE``.640 641 Other config options are also scaled accordingly:642 * training steps and warmup steps are scaled inverse proportionally.643 * learning rate are scaled proportionally, following :paper:`ImageNet in 1h`.644 645 For example, with the original config like the following:646 647 .. code-block:: yaml648 649 IMS_PER_BATCH: 16650 BASE_LR: 0.1651 REFERENCE_WORLD_SIZE: 8652 MAX_ITER: 5000653 STEPS: (4000,)654 CHECKPOINT_PERIOD: 1000655 656 When this config is used on 16 GPUs instead of the reference number 8,657 calling this method will return a new config with:658 659 .. code-block:: yaml660 661 IMS_PER_BATCH: 32662 BASE_LR: 0.2663 REFERENCE_WORLD_SIZE: 16664 MAX_ITER: 2500665 STEPS: (2000,)666 CHECKPOINT_PERIOD: 500667 668 Note that both the original config and this new config can be trained on 16 GPUs.669 It's up to user whether to enable this feature (by setting ``REFERENCE_WORLD_SIZE``).670 671 Returns:672 CfgNode: a new config. Same as original if ``cfg.SOLVER.REFERENCE_WORLD_SIZE==0``.673 """674 old_world_size = cfg.SOLVER.REFERENCE_WORLD_SIZE675 if old_world_size == 0 or old_world_size == num_workers:676 return cfg677 cfg = cfg.clone()678 frozen = cfg.is_frozen()679 cfg.defrost()680 681 assert (682 cfg.SOLVER.IMS_PER_BATCH % old_world_size == 0683 ), "Invalid REFERENCE_WORLD_SIZE in config!"684 scale = num_workers / old_world_size685 bs = cfg.SOLVER.IMS_PER_BATCH = int(round(cfg.SOLVER.IMS_PER_BATCH * scale))686 lr = cfg.SOLVER.BASE_LR = cfg.SOLVER.BASE_LR * scale687 max_iter = cfg.SOLVER.MAX_ITER = int(round(cfg.SOLVER.MAX_ITER / scale))688 warmup_iter = cfg.SOLVER.WARMUP_ITERS = int(round(cfg.SOLVER.WARMUP_ITERS / scale))689 cfg.SOLVER.STEPS = tuple(int(round(s / scale)) for s in cfg.SOLVER.STEPS)690 cfg.TEST.EVAL_PERIOD = int(round(cfg.TEST.EVAL_PERIOD / scale))691 cfg.SOLVER.CHECKPOINT_PERIOD = int(round(cfg.SOLVER.CHECKPOINT_PERIOD / scale))692 cfg.SOLVER.REFERENCE_WORLD_SIZE = num_workers # maintain invariant693 logger = logging.getLogger(__name__)694 logger.info(695 f"Auto-scaling the config to batch_size={bs}, learning_rate={lr}, "696 f"max_iter={max_iter}, warmup={warmup_iter}."697 )698 699 if frozen:700 cfg.freeze()701 return cfg702 703 704# Access basic attributes from the underlying trainer705for _attr in ["model", "data_loader", "optimizer"]:706 setattr(707 DefaultTrainer,708 _attr,709 property(710 # getter711 lambda self, x=_attr: getattr(self._trainer, x),712 # setter713 lambda self, value, x=_attr: setattr(self._trainer, x, value),714 ),715 )716 