all_code_base / lrm /flux /trainer /tasks /base_task.py
aryadomain's picture
Add files using upload-large-folder tool
533920b verified
Raw
History Blame Contribute Delete
2.78 kB
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