"""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