import logging import os import json try: import wandb except ImportError: wandb = None def setup_logging(log_file, level, include_host=False): if include_host: import socket hostname = socket.gethostname() formatter = logging.Formatter( f'%(asctime)s | {hostname} | %(levelname)s | %(message)s', datefmt='%Y-%m-%d,%H:%M:%S') else: formatter = logging.Formatter('%(asctime)s | %(levelname)s | %(message)s', datefmt='%Y-%m-%d,%H:%M:%S') logging.root.setLevel(level) loggers = [logging.getLogger(name) for name in logging.root.manager.loggerDict] for logger in loggers: logger.setLevel(level) stream_handler = logging.StreamHandler() stream_handler.setFormatter(formatter) logging.root.addHandler(stream_handler) if log_file: file_handler = logging.FileHandler(filename=log_file) file_handler.setFormatter(formatter) logging.root.addHandler(file_handler) def write_eval_log(args, log_data, data, epoch, metrics, tb_writer=None): if args.save_logs: if tb_writer is not None: for name, val in log_data.items(): tb_writer.add_scalar(name, val, epoch) with open(os.path.join(args.checkpoint_path, "results.jsonl"), "a+") as f: f.write(json.dumps(metrics)) f.write("\n") if args.wandb: assert wandb is not None, 'Please install wandb.' if 'train' in data: dataloader = data['train'].dataloader num_batches_per_epoch = dataloader.num_batches // args.accum_freq step = num_batches_per_epoch * epoch else: step = None log_data['epoch'] = epoch wandb.log(log_data, step=step)