| 49 | |
| 50 | |
| 51 | def update_batch_size_info(cfg: DictConfig): |
| 52 | global_batch_size, device_microbatch_size = cfg.global_train_batch_size, cfg.device_train_microbatch_size |
| 53 | if global_batch_size % dist.get_world_size() != 0: |
| 54 | raise ValueError( |
| 55 | f"Global batch size {global_batch_size} is not divisible by {dist.get_world_size()} " |
| 56 | "as a result, the batch size would be truncated, please adjust `global_batch_size` " |
| 57 | f"to be divisible by world size, {dist.get_world_size()}." |
| 58 | ) |
| 59 | device_train_batch_size = global_batch_size // dist.get_world_size() |
| 60 | if isinstance(device_microbatch_size, int): |
| 61 | if device_microbatch_size > device_train_batch_size: |
| 62 | print( |
| 63 | f"WARNING: device_train_microbatch_size > device_train_batch_size, " |
| 64 | f"will be reduced from {device_microbatch_size} -> {device_train_batch_size}." |
| 65 | ) |
| 66 | device_microbatch_size = device_train_batch_size |
| 67 | cfg.n_gpus = dist.get_world_size() |
| 68 | cfg.device_train_batch_size = device_train_batch_size |
| 69 | cfg.device_train_microbatch_size = device_microbatch_size |
| 70 | |
| 71 | # Safely set `device_eval_microbatch_size` if not provided by user |
| 72 | if "device_eval_microbatch_size" not in cfg: |
| 73 | if cfg.device_train_microbatch_size == "auto": |
| 74 | cfg.device_eval_microbatch_size = 1 |
| 75 | else: |
| 76 | cfg.device_eval_microbatch_size = cfg.device_train_microbatch_size |
| 77 | |
| 78 | global_eval_batch_size, device_eval_microbatch_size = ( |
| 79 | cfg.get("global_eval_batch_size", global_batch_size), |
| 80 | cfg.device_eval_microbatch_size, |
| 81 | ) |
| 82 | device_eval_batch_size = global_eval_batch_size // dist.get_world_size() |
| 83 | if isinstance(device_eval_microbatch_size, int): |
| 84 | if device_eval_microbatch_size > device_eval_microbatch_size: |
| 85 | print( |
| 86 | f"WARNING: device_eval_microbatch_size > device_eval_batch_size, " |
| 87 | f"will be reduced from {device_eval_microbatch_size} -> {device_eval_batch_size}." |
| 88 | ) |
| 89 | device_eval_microbatch_size = device_eval_batch_size |
| 90 | cfg.device_eval_batch_size = device_eval_batch_size |
| 91 | cfg.device_eval_microbatch_size = device_eval_microbatch_size |
| 92 | return cfg |
| 93 | |
| 94 | |
| 95 | # from timm: https://github.com/huggingface/pytorch-image-models/blob/main/timm/optim/optim_factory.py |