TSRDA / main_method /code /training_feature_clusters.py
Dhruv1000's picture
Organize complete final models, all ablations, logs and checkpoints with visual guides (part 7)
71d64bb verified
Raw History Blame Contribute Delete
3.33 kB
"""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