File size: 3,329 Bytes
71d64bb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Prototype (three training-feature cluster centres per class) readout and metrics. Extracted unchanged from the scored implementation."""
import numpy as np
import torch
import torch.nn.functional as F
from temporal_cnn import valid_from_index_positions

def metrics_cm(cm):
    tp = np.diag(cm)[1:19].astype(float)
    gt, pp = cm.sum(1)[1:19], cm.sum(0)[1:19]
    iou = tp / np.maximum(gt + pp - tp, 1)
    f1 = 2 * tp / np.maximum(gt + pp, 1)
    total = cm.sum()
    oa = cm.trace() / max(total, 1)
    chance = np.dot(cm.sum(0).astype(float), cm.sum(1)) / max(total ** 2, 1)
    return {'miou': float(iou.mean()), 'oa': float(oa),
            'macro_f1': float(f1.mean()), 'kappa': float((oa-chance)/max(1-chance,1e-12)),
            'per_class_iou': iou.tolist(), 'confusion_matrix': cm.tolist()}


class DepthCollector:
    def __init__(self, encoder):
        self.handles = [block.register_forward_hook(self.hook(i))
                        for i, block in enumerate(encoder.temporal_encoder.blocks)]

    def begin(self, x, pos):
        b, _, _, h, w = x.shape
        self.shape = b, h, w
        self.valid = valid_from_index_positions(pos)[:, None, None]
        self.valid = self.valid.expand(b, h, w, -1).reshape(b*h*w, -1)
        self.parts = [[], [], []]
        self.offsets = [0, 0, 0]

    def hook(self, depth):
        def collect(module, inputs, output):
            start = self.offsets[depth]
            valid = self.valid[start:start + output.shape[0]]
            if valid.shape[1] < output.shape[1]:
                valid = F.pad(valid, (0, output.shape[1]-valid.shape[1]), value=False)
            weights = valid[..., None].float()
            values = output.float()
            count = weights.sum(1).clamp_min(1)
            mean = (values * weights).sum(1) / count
            var = ((values - mean[:, None]).square() * weights).sum(1) / count
            summary = torch.cat((mean, var.clamp_min(0).sqrt()), -1)
            self.parts[depth].append(summary.detach().half().cpu().numpy())
            self.offsets[depth] += len(output)
        return collect

    def end(self):
        assert self.offsets == [np.prod(self.shape)] * 3
        return np.stack([np.concatenate(parts) for parts in self.parts], axis=1)

    def close(self):
        for handle in self.handles:
            handle.remove()


def rows(features, indices, depth):
    data = np.asarray(features[indices])
    return (data[:,depth,:] if depth>=0 else data.reshape(len(data),-1)).astype(np.float32)


def normalized(features, indices, depth, mean, std):
    return (torch.from_numpy(rows(features,indices,depth)).cuda()-mean)/std


def probability(features, depth, mean, std, kind, state):
    out=np.zeros((len(features),20),np.float32)
    with torch.inference_mode():
        for start in range(0,len(features),4096):
            ix=np.arange(start,min(start+4096,len(features)))
            x=normalized(features,ix,depth,mean,std)
            if kind=='linear':
                logits=F.linear(x,state['weight'],state['bias'])
            else:
                embeddings=F.normalize(x,dim=-1)
                logits=torch.einsum('nd,ckd->nck',embeddings,state)/.1
                logits=logits.logsumexp(-1)-np.log(3)
            out[start:start+len(ix),1:19]=logits.softmax(-1).cpu().numpy()
    return out