()
| 316 | |
| 317 | |
| 318 | def 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()) |
no test coverage detected