File size: 3,026 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
"""Per-pixel spectral time-series CNNs: original mean/max and attention/max variants."""
import torch
from torch import nn

def valid_from_index_positions(pos):
    """The unmodified dataset uses 0,1,...,T-1 and pads positions with zero.

Temporal subsampling may drop acquisition zero; its first retained position is
then positive. Only an initial zero is a valid acquisition-zero token.
This cannot be reused for arbitrary calendar positions without a length mask.
"""
    first_is_zero = pos[:, :1].eq(0)
    first = torch.zeros_like(pos, dtype=torch.bool)
    first[:, :1] = first_is_zero
    valid = pos.gt(0) | first
    if not valid.any(dim=1).all():
        raise ValueError('Each time series must contain an observation')
    return valid


class TemporalCNN(nn.Module):
    def __init__(self, classes=18):
        super().__init__()
        self.input=nn.Linear(10,64)
        self.norms=nn.ModuleList([nn.LayerNorm(64) for _ in range(4)])
        self.convs=nn.ModuleList([nn.Conv1d(64,64,3,padding=d,dilation=d,groups=64) for d in [1,2,4,8]])
        self.mix=nn.ModuleList([nn.Linear(64,64) for _ in range(4)])
        self.head=nn.Sequential(nn.Linear(128,64),nn.GELU(),nn.Linear(64,classes))

    def forward(self,x,pos):
        b,t,c,h,w=x.shape
        valid=valid_from_index_positions(pos)[:,None,None].expand(b,h,w,t).reshape(-1,t,1)
        flat=x.permute(0,3,4,1,2).reshape(-1,t,c)
        outputs=[]
        for start in range(0,len(flat),2048):
            mask=valid[start:start+2048];z=self.input(flat[start:start+2048])*mask
            for norm,conv,mix in zip(self.norms,self.convs,self.mix):
                z=(z+mix(torch.nn.functional.gelu(conv(norm(z).transpose(1,2)).transpose(1,2))))*mask
            mean=z.sum(1)/mask.sum(1).clamp_min(1)
            maximum=z.masked_fill(~mask,torch.finfo(z.dtype).min).amax(1)
            outputs.append(self.head(torch.cat([mean,maximum],-1)))
        return torch.cat(outputs).reshape(b,h,w,-1).permute(0,3,1,2)


class AttentionCNN(TemporalCNN):
    def __init__(self,classes=18):
        super().__init__(classes);self.attention=nn.Linear(64,1)
    def forward(self,x,pos):
        b,t,c,h,w=x.shape;valid=valid_from_index_positions(pos)[:,None,None].expand(b,h,w,t).reshape(-1,t,1);flat=x.permute(0,3,4,1,2).reshape(-1,t,c);out=[]
        for start in range(0,len(flat),2048):
            mask=valid[start:start+2048];z=self.input(flat[start:start+2048])*mask
            for norm,conv,mix in zip(self.norms,self.convs,self.mix):z=(z+mix(torch.nn.functional.gelu(conv(norm(z).transpose(1,2)).transpose(1,2))))*mask
            score=self.attention(z).masked_fill(~mask,torch.finfo(z.dtype).min)
            pooled=(score.softmax(1)*z).sum(1);maximum=z.masked_fill(~mask,torch.finfo(z.dtype).min).amax(1)
            out.append(self.head(torch.cat([pooled,maximum],-1)))
        return torch.cat(out).reshape(b,h,w,-1).permute(0,3,1,2)


def make_model(kind,classes=18):
    return {"cnn_control":TemporalCNN,"attention":AttentionCNN}[kind](classes)