MCPcopy Create free account
hub / github.com/LUMIA-Group/MemoryDecoder / main

Function main

train_memdec.py:318–773  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

316
317
318def main():
319 args = parse_args()
320
321 accelerator_log_kwargs = {}
322
323 if args.with_tracking:
324 accelerator_log_kwargs["log_with"] = args.report_to
325 accelerator_log_kwargs["project_dir"] = args.output_dir
326
327 accelerator = Accelerator(gradient_accumulation_steps=args.gradient_accumulation_steps, **accelerator_log_kwargs)
328
329 if args.report_to == "wandb":
330 accelerator.init_trackers(
331 project_name=args.project_name,
332 config=args,
333 init_kwargs={
334 "wandb": {
335 "name": args.run_name if args.run_name is not None else None,
336 "group": args.group_name if args.group_name is not None else None,
337 "save_code": True,
338 },
339 }
340 )
341
342 # Make one log on every process with the configuration for debugging.
343 if accelerator.is_local_main_process:
344 datasets.utils.logging.set_verbosity_warning()
345 transformers.utils.logging.set_verbosity_info()
346 log_level = logging.INFO
347 else:
348 datasets.utils.logging.set_verbosity_error()
349 transformers.utils.logging.set_verbosity_error()
350 log_level = logging.ERROR
351
352 logger.remove()
353 logger.add(sys.stdout, format = "<green>{time:YYYY-MM-DD HH:mm:ss}</green> | <level>{level}</level> | <blue>{process.name}</blue> | <cyan>{name}</cyan>:<cyan>{function}</cyan>:<cyan>{line}</cyan> - <level>{message}</level>", level=log_level)
354
355 # Intercept default logging and transform to loguru
356 class InterceptHandler(logging.Handler):
357 def emit(self, record: logging.LogRecord) -> None:
358 # Get corresponding Loguru level if it exists.
359 level: str | int
360 try:
361 level = logger.level(record.levelname).name
362 except ValueError:
363 level = record.levelno
364
365 # Find caller from where originated the logged message.
366 frame, depth = inspect.currentframe(), 0
367 while frame and (depth == 0 or frame.f_code.co_filename == logging.__file__):
368 frame = frame.f_back
369 depth += 1
370
371 logger.opt(depth=depth, exception=record.exc_info).log(level, record.getMessage())
372
373 logging.basicConfig(handlers=[InterceptHandler()], level=log_level, force=True)
374 transformers.utils.logging.disable_default_handler()
375 transformers.utils.logging.add_handler(InterceptHandler())

Callers 1

train_memdec.pyFile · 0.70

Calls 4

kl_loss_evaluateFunction · 0.90
kl_loss_tokenFunction · 0.90
parse_argsFunction · 0.70
InterceptHandlerClass · 0.70

Tested by

no test coverage detected