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