NowcastNet_Earth / model /nowcastnet.py
yzt15806542928's picture
Upload folder using huggingface_hub
439c523 verified
Raw
History Blame Contribute Delete
3.49 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
import numpy as np
from .layers_utils import warp, make_grid
from .generative_network import Generative_Encoder, Generative_Decoder
from .evolution_network import Evolution_Network
from .noise_projector import Noise_Projector
class Net(nn.Module):
def __init__(self, configs):
super(Net, self).__init__()
self.configs = configs
self.pred_length = self.configs.total_length - self.configs.input_length
self.evo_net = Evolution_Network(self.configs.input_length, self.pred_length, base_c=32)
self.gen_enc = Generative_Encoder(self.configs.total_length, base_c=self.configs.ngf)
self.gen_dec = Generative_Decoder(self.configs)
self.proj = Noise_Projector(self.configs.ngf, configs)
sample_tensor = torch.zeros(1, 1, self.configs.img_height, self.configs.img_width)
self.register_buffer("grid", make_grid(sample_tensor), persistent=False)
def forward(self, all_frames):
if all_frames.ndim == 4:
all_frames = all_frames.unsqueeze(-1)
if all_frames.ndim != 5 or all_frames.shape[-1] < 1:
raise ValueError("Expected frames with shape [B,T,H,W,C], C>=1")
all_frames = all_frames[:, :, :, :, :1]
frames = all_frames.permute(0, 1, 4, 2, 3)
batch = frames.shape[0]
height = frames.shape[3]
width = frames.shape[4]
# Input Frames
input_frames = frames[:, :self.configs.input_length]
input_frames = input_frames.reshape(batch, self.configs.input_length, height, width)
# Evolution Network
intensity, motion = self.evo_net(input_frames)
motion_ = motion.reshape(batch, self.pred_length, 2, height, width)
intensity_ = intensity.reshape(batch, self.pred_length, 1, height, width)
series = []
last_frames = all_frames[:, (self.configs.input_length - 1):self.configs.input_length, :, :, 0]
grid = self.grid.to(frames.device).repeat(batch, 1, 1, 1)
for i in range(self.pred_length):
last_frames = warp(last_frames, motion_[:, i], grid, mode="nearest", padding_mode="border")
last_frames = last_frames + intensity_[:, i]
series.append(last_frames)
evo_result = torch.cat(series, dim=1)
evo_result = evo_result/128
# Generative Network
evo_feature = self.gen_enc(torch.cat([input_frames, evo_result], dim=1))
noise = torch.randn(batch, self.configs.ngf, max(1, height // 32), max(1, width // 32), device=frames.device)
projected = self.proj(noise)
# The official reshape/permute is a 4x pixel shuffle: 32*ngf
# projected channels become 2*ngf channels at H/8 x W/8. Expressing
# it directly preserves the checkpoint layout and also works at the
# official 512x512 resolution (where the released literal reshape is
# dimensionally inconsistent).
if projected.shape[2] * 4 == height // 8 and projected.shape[3] * 4 == width // 8:
noise_feature = F.pixel_shuffle(projected, 4)
else:
target_hw = (max(1, height // 8), max(1, width // 8))
noise_feature = F.interpolate(projected[:, : 2 * self.configs.ngf], size=target_hw, mode="nearest")
feature = torch.cat([evo_feature, noise_feature], dim=1)
gen_result = self.gen_dec(feature, evo_result)
return gen_result.unsqueeze(-1)