MCPcopy Create free account
hub / github.com/GanjinZero/BioBART / construct_arguments

Function construct_arguments

pretrain_src/train.py:136–194  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

134
135
136def 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

Callers 1

mainFunction · 0.85

Calls 3

LoggerClass · 0.90
get_argumentsFunction · 0.85
loadMethod · 0.45

Tested by

no test coverage detected