()
| 134 | |
| 135 | |
| 136 | def construct_arguments(): |
| 137 | |
| 138 | args = get_arguments() |
| 139 | |
| 140 | # Prepare Logger |
| 141 | logger = Logger(cuda=torch.cuda.is_available() and not args.no_cuda) |
| 142 | args.logger = logger |
| 143 | config = json.load(open(args.config_file, 'r', encoding='utf-8')) |
| 144 | |
| 145 | # choose dataset and training config based on the given sequence length |
| 146 | seq_len = str(args.max_seq_length) |
| 147 | |
| 148 | datasets = config["data"]["mixed_seq_datasets"][seq_len] |
| 149 | del config["data"]["mixed_seq_datasets"] |
| 150 | training = config["mixed_seq_training"][seq_len] |
| 151 | del config["mixed_seq_training"] |
| 152 | config["data"]["datasets"] = datasets |
| 153 | config["training"] = training |
| 154 | args.config = config |
| 155 | |
| 156 | args.max_steps = args.config["training"]["total_training_steps"] |
| 157 | |
| 158 | args.job_name = config['name'] if args.job_name is None else args.job_name |
| 159 | print("Running Config File: ", args.job_name) |
| 160 | # Setting the distributed variables |
| 161 | print("Args = {}".format(args)) |
| 162 | |
| 163 | # Setting all the seeds so that the task is random but same accross processes |
| 164 | random.seed(args.seed) |
| 165 | np.random.seed(args.seed) |
| 166 | torch.manual_seed(args.seed) |
| 167 | torch.cuda.manual_seed_all(args.seed) |
| 168 | |
| 169 | os.makedirs(args.output_dir, exist_ok=True) |
| 170 | args.saved_model_path = os.path.join(args.output_dir, "saved_models/", |
| 171 | args.job_name) |
| 172 | |
| 173 | # args.n_gpu = 1 |
| 174 | |
| 175 | tokenizer = BartTokenizer.from_pretrained(config["bart_token_file"]) |
| 176 | args.tokenizer = tokenizer |
| 177 | |
| 178 | # Set validation dataset path |
| 179 | if args.validation_data_path_prefix is None: |
| 180 | logging.warning( |
| 181 | 'Skipping validation because validation_data_path_prefix is unspecified' |
| 182 | ) |
| 183 | |
| 184 | # Issue warning if early exit from epoch is configured |
| 185 | if args.max_steps < sys.maxsize: |
| 186 | logging.warning( |
| 187 | 'Early training exit is set after {} global steps'.format( |
| 188 | args.max_steps)) |
| 189 | |
| 190 | if args.max_steps_per_epoch < sys.maxsize: |
| 191 | logging.warning('Early epoch exit is set after {} global steps'.format( |
| 192 | args.max_steps_per_epoch)) |
| 193 |
no test coverage detected