Parse all arguments.
(extra_args_provider=None, defaults={}, ignore_unknown_args=False)
| 23 | |
| 24 | |
| 25 | def parse_args(extra_args_provider=None, defaults={}, ignore_unknown_args=False): |
| 26 | """Parse all arguments.""" |
| 27 | parser = argparse.ArgumentParser( |
| 28 | description="Megatron-LM Arguments", allow_abbrev=False |
| 29 | ) |
| 30 | |
| 31 | # Standard arguments. |
| 32 | parser = _add_network_size_args(parser) |
| 33 | parser = _add_regularization_args(parser) |
| 34 | parser = _add_training_args(parser) |
| 35 | parser = _add_initialization_args(parser) |
| 36 | parser = _add_learning_rate_args(parser) |
| 37 | parser = _add_checkpointing_args(parser) |
| 38 | parser = _add_mixed_precision_args(parser) |
| 39 | parser = _add_distributed_args(parser) |
| 40 | parser = _add_validation_args(parser) |
| 41 | parser = _add_data_args(parser) |
| 42 | parser = _add_autoresume_args(parser) |
| 43 | parser = _add_biencoder_args(parser) |
| 44 | parser = _add_vit_args(parser) |
| 45 | parser = _add_logging_args(parser) |
| 46 | parser = _add_zero_args(parser) |
| 47 | parser = _add_memoryopt_args(parser) |
| 48 | parser = _add_activation_checkpoint_args(parser) |
| 49 | parser = _add_inference_args(parser) |
| 50 | |
| 51 | # Custom arguments. |
| 52 | if extra_args_provider is not None: |
| 53 | parser = extra_args_provider(parser) |
| 54 | |
| 55 | parser = deepspeed.add_config_arguments(parser) |
| 56 | |
| 57 | # Parse. |
| 58 | if ignore_unknown_args: |
| 59 | args, _ = parser.parse_known_args() |
| 60 | else: |
| 61 | args = parser.parse_args() |
| 62 | |
| 63 | # helper argument to set deepspeed pipeline parallel or not |
| 64 | args.ds_pipeline_enabled = not args.no_pipeline_parallel |
| 65 | |
| 66 | # Distributed args. |
| 67 | args.rank = int(os.getenv("RANK", "0")) |
| 68 | args.world_size = int(os.getenv("WORLD_SIZE", "1")) |
| 69 | # Tensor model parallel size. |
| 70 | args.tensor_model_parallel_size = min( |
| 71 | args.tensor_model_parallel_size, args.world_size |
| 72 | ) |
| 73 | assert ( |
| 74 | args.world_size % args.tensor_model_parallel_size == 0 |
| 75 | ), "world size" " ({}) is not divisible by tensor model parallel size ({})".format( |
| 76 | args.world_size, args.tensor_model_parallel_size |
| 77 | ) |
| 78 | # Pipeline model parallel size. |
| 79 | args.pipeline_model_parallel_size = min( |
| 80 | args.pipeline_model_parallel_size, |
| 81 | (args.world_size // args.tensor_model_parallel_size), |
| 82 | ) |
no test coverage detected