File size: 7,151 Bytes
07fcdfe
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
from copy import deepcopy
import string
from time import perf_counter
from typing import Callable

from lightning.fabric.wrappers import _FabricModule
import numpy as np
import torch
import torch.nn as nn


def _test_condition(condition: bool, text: str):
    if not condition:
        raise ValueError(text)


# SAMPLE WEIGHTS
""" def cal_weights_model_1_forward(dataset):
    #predicting 'diseased'
    labels = []
    for data in dataset:
        labels += data.diseased.tolist()
    labels_tensor = torch.tensor(labels).squeeze()
    n_positive = labels_tensor.nonzero().size(0)
    n_negative = labels_tensor.size(0) - n_positive
    n_full = labels_tensor.size(0)
    return torch.tensor([n_full / (2 * n_negative), n_full / (2 * n_positive)])

def cal_weights_model_1_backward(dataset):
    #predicting 'treated'
    labels = []
    for data in dataset:
        labels += data.treated.tolist()
    labels_tensor = torch.tensor(labels).squeeze()
    n_positive = labels_tensor.nonzero().size(0)
    n_negative = labels_tensor.size(0) - n_positive
    n_full = labels_tensor.size(0)
    return torch.tensor([n_full / (2 * n_negative), n_full / (2 * n_positive)])

def cal_weights_model_2_backward(dataset):
    #predicting 'intervention'
    labels = []
    for data in dataset:
        labels += data.intervention.tolist()
    labels_tensor = torch.tensor(labels).squeeze()
    n_positive = labels_tensor.nonzero().size(0)
    n_negative = labels_tensor.size(0) - n_positive
    n_full = labels_tensor.size(0)
    return torch.tensor([n_full / (2 * n_negative), n_full / (2 * n_positive)]) """


def calculate_loss_sample_weights(dataset, kind: str) -> torch.Tensor:
    _test_condition(kind in {"diseased", "treated", "intervention"}, "`kind` should be one of (diseased, treated, intervention)")
    labels = []
    for data in dataset:
        labels += getattr(data, kind).tolist()
    labels_tensor = torch.tensor(labels).squeeze()
    n_positive = labels_tensor.nonzero().size(0)
    n_negative = labels_tensor.size(0) - n_positive
    n_full = labels_tensor.size(0)
    return torch.tensor([n_full/(2*n_negative), n_full/(2*n_positive)])


""" def get_threshold_healthy(dataset):
    all_healthy_values = []
    for data in dataset:
        all_healthy_values.append(data.healthy.cpu())
    percentiles = torch.Tensor(np.percentile(torch.stack(all_healthy_values).flatten(), [e for e in np.arange(0,100,0.2)] + [100]))
    return percentiles

def get_threshold_diseased(dataset):
    all_diseased_values = []
    for data in dataset:
        all_diseased_values.append(data.diseased.cpu())
    percentiles = torch.Tensor(np.percentile(torch.stack(all_diseased_values).flatten(), [e for e in np.arange(0,100,0.2)] + [100]))
    return percentiles

def get_threshold_treated(dataset):
    all_treated_values = []
    for data in dataset:
        all_treated_values.append(data.treated.cpu())
    percentiles = torch.Tensor(np.percentile(torch.stack(all_treated_values).flatten(), np.arange(0, 100.2, 0.2)))
    return percentiles """


def _get_thresholds(dataset, kind: str):
    _test_condition(kind in {"healthy", "diseased", "treated"}, "`kind` should be one of (diseased, treated, healthy)")
    all_values = [getattr(data, kind).cpu() for data in dataset]
    percentiles = torch.tensor(np.percentile(torch.stack(all_values).flatten(), [e for e in np.arange(0, 100, 0.2)] + [100]))
    return percentiles


def get_thresholds(dataset):
    return {
        'healthy': _get_thresholds(dataset.train_dataset_forward, "healthy") if hasattr(dataset, 'train_dataset_forward') else None,
        'diseased': _get_thresholds(dataset.train_dataset_backward, "diseased"),
        'treated': _get_thresholds(dataset.train_dataset_backward, "treated")
    }


class EarlyStopping:

    def __init__(self, patience: int = 15, skip: int = 0, minmax: str = "min", rope: float = 1e-5,
                 model: _FabricModule = None, save_path: str = None):

        self.skip = skip
        self.patience = patience
        self.rope = abs(rope)

        self.minmax = minmax
        self.comparison_f = (lambda x, y: x < y-self.rope) if self.minmax == "min" else (lambda x, y: x > y+self.rope)

        self.reset()

        self.model = model
        self.save_path = save_path

        self.successful_comparison = (self._save_model if (self.save_path and self.model) else lambda: None)
        self.load_model = (self._load_model if (self.save_path and self.model) else lambda: None)

    def _save_model(self):
        tmp_model = deepcopy(self.model.module)
        torch.save({"epoch": self.skip_counter, "model_state_dict": tmp_model.cpu().state_dict()}, self.save_path)

    def _load_model(self) -> nn.Module:
        checkpoint = torch.load(self.save_path)
        tmp_model = deepcopy(self.model.module)
        tmp_model.load_state_dict(checkpoint["model_state_dict"])
        return tmp_model

    def reset(self):
        self.counter = 0
        self.skip_counter = 0
        self.is_stopped = False
        self.value = float("inf") if self.minmax == "min" else -float("inf")

    def __call__(self, value):
        self.skip_counter += 1
        if self.skip_counter < self.skip:
            if self.comparison_f(value, self.value): # even when skipping, save best value
                self.value = value
                self.successful_comparison()
            return False

        if self.comparison_f(value, self.value):
            self.value = value
            self.counter = 0
            self.successful_comparison()
        else:
            self.counter += 1
            if self.counter >= self.patience:
                self.is_stopped = True
                return True

        return False


class DummyEarlyStopping(EarlyStopping):

    def __init__(self, patience: int = 15, skip: int = 0, minmax: str = "min", rope: float = 1e-5,
                 model: _FabricModule = None, save_path: str = None):
        super().__init__(patience, skip, minmax, rope, model, save_path)

        self.successful_comparison = lambda: None
        self.load_model = lambda: None

    def __call__(self, value):
        return False


def tictoc(*args):
    # https://stackoverflow.com/questions/3931627/how-to-build-a-decorator-with-optional-parameters
    def wrap(function: Callable):
        def wrapped_f(*args, **kwargs):
            tic = perf_counter()
            result = function(*args, **kwargs)
            toc = perf_counter()
            print(text_to_format.format(toc-tic))
            return result
        return wrapped_f
    if len(args) >= 1 and callable(args[0]):
        text_to_format: str = args[1] if len(args) >= 2 else "{}secs"
        to_return = wrap(args[0])
    else:
        text_to_format = args[0] if args else "{}secs"
        to_return = wrap
    matches = [tup[1] for tup in string.Formatter().parse(text_to_format) if tup[1] is not None]
    if len(matches) != 1:
        raise ValueError(r"tictoc decorator requires string with one {}!")
    return to_return


class DummyWriter:
    def __init__(self, *args, **kwargs):
        pass

    def add_scalar(self, *args, **kwargs):
        pass