from dataclasses import dataclass import torch from PIL import Image from accelerate.logging import get_logger from accelerate.utils import LoggerType logger = get_logger(__name__) def flatten(list_of_lists): return [item for sublist in list_of_lists for item in sublist] @dataclass class BaseTaskConfig: limit_examples_to_wandb: int = 50 pass class BaseTask: def __init__(self, cfg: BaseTaskConfig, accelerator): self.accelerator = accelerator self.cfg = cfg def train_step(self, model, criterion, batch): pass def valid_step(self, model, criterion, batch): pass def evaluate(self, model, criterion, dataloader): pass def log_to_wandb(self, eval_dict, table_name="test_predictions"): if not self.accelerator.is_main_process or not LoggerType.WANDB == self.accelerator.cfg.log_with: logger.info("Not logging to wandb") return import wandb logger.info("Uploading to wandb") for key, value in eval_dict.items(): eval_dict[key] = [wandb.Image(maybe_img) if isinstance(maybe_img, Image.Image) else maybe_img for maybe_img in value] if self.cfg.limit_examples_to_wandb > 0: eval_dict[key] = eval_dict[key][:self.cfg.limit_examples_to_wandb] columns, predictions = list(zip(*sorted(eval_dict.items()))) predictions += ([self.accelerator.global_step] * len(predictions[0]),) columns += ("global_step",) data = list(zip(*predictions)) table = wandb.Table(columns=list(columns), data=data) wandb.log({table_name: table}, commit=False, step=self.accelerator.global_step) @staticmethod def gather_iterable(it, num_processes): if num_processes <= 1: return it if not torch.distributed.is_available() or not torch.distributed.is_initialized(): return it output_objects = [None for _ in range(num_processes)] torch.distributed.all_gather_object(output_objects, it) return flatten(output_objects) @torch.no_grad() def valid_step(self, model, criterion, batch): loss = criterion(model, batch) return loss def gather_dict(self, eval_dict): if self.accelerator.num_processes <= 1: return eval_dict if not torch.distributed.is_available() or not torch.distributed.is_initialized(): logger.warning("Distributed process group is not initialized; skipping gather.") return eval_dict logger.info("Gathering dict from all processes...") for k, v in eval_dict.items(): eval_dict[k] = self.gather_iterable(v, self.accelerator.num_processes) return eval_dict