MCPcopy Create free account
hub / github.com/SkyworkAI/Skywork / main

Function main

train/train.py:221–451  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

219
220
221def main():
222
223 parser = transformers.HfArgumentParser(
224 (ModelArguments, DataArguments, TrainingArguments, LoraArguments)
225 )
226
227 (
228 model_args,
229 data_args,
230 training_args,
231 lora_args,
232 ) = parser.parse_args_into_dataclasses()
233
234 device_map = "auto"
235 world_size = int(os.environ.get("WORLD_SIZE", 1))
236 ddp = world_size != 1
237 if lora_args.use_q_lora:
238 device_map = {"": int(os.environ.get("LOCAL_RANK") or 0)} if ddp else None
239 if len(training_args.fsdp) > 0 or deepspeed.is_deepspeed_zero3_enabled():
240 logging.warning(
241 "FSDP or ZeRO3 are not incompatible with QLoRA."
242 )
243
244 send_example_telemetry("run_pt", model_args, data_args)
245
246 # Setup logging
247 logging.basicConfig(format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",datefmt="%m/%d/%Y %H:%M:%S",
248 level=logging.INFO, # if training_args.local_rank in [-1, 0] else logging.WARN,
249 handlers=[logging.StreamHandler(sys.stdout)],)
250
251
252 if training_args.should_log:
253 # The default of training_args.log_level is passive, so we set log level at info here to have that default.
254 transformers.utils.logging.set_verbosity_info()
255
256 log_level = training_args.get_process_log_level()
257 logger.setLevel(log_level)
258 datasets.utils.logging.set_verbosity(log_level)
259 transformers.utils.logging.set_verbosity(log_level)
260 transformers.utils.logging.enable_default_handler()
261 transformers.utils.logging.enable_explicit_format()
262
263 # Log on each process the small summary:
264 logger.warning(
265 f"Process rank: {training_args.local_rank}, device: {training_args.device}, n_gpu: {training_args.n_gpu}"
266 + f"distributed training: {bool(training_args.local_rank != -1)}, 16-bits training: {training_args.fp16 or training_args.bf16}"
267 )
268
269 # Detecting last checkpoint.
270 last_checkpoint = None
271 if os.path.isdir(training_args.output_dir) and training_args.do_train and not training_args.overwrite_output_dir:
272 last_checkpoint = get_last_checkpoint(training_args.output_dir)
273 if last_checkpoint is None and len(os.listdir(training_args.output_dir)) > 0:
274 raise ValueError(
275 f"Output directory ({training_args.output_dir}) already exists and is not empty. "
276 "Use --overwrite_output_dir to overcome."
277 )
278 elif last_checkpoint is not None and training_args.resume_from_checkpoint is None:

Callers 1

train.pyFile · 0.70

Calls 4

set_seedFunction · 0.85

Tested by

no test coverage detected