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