Download main_method/code/temporal_cnn.py from Dhruv1000/TSRDA: direct link, hf CLI and curl.
- Browser
- Download file 3.03 kB
-
https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/temporal_cnn.py
- Command line
-
hf download hf://Dhruv1000/TSRDA/main_method/code/temporal_cnn.py
-
curl -L -o temporal_cnn.py https://huggingface.co/Dhruv1000/TSRDA/resolve/main/main_method/code/temporal_cnn.py
3.03 kB
| """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) | |