File size: 2,783 Bytes
533920b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
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