File size: 4,478 Bytes
338c3e4 | 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 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 | from typing import Optional, List
import torch
import torch.nn as nn
from torch import Tensor
from .base_model import AutoCfdModel
from onescience.modules.decoder.unet_decoder import UNetDecoder2D
from onescience.modules.encoder.unet_encoder import UNetEncoder2D
from onescience.modules.head.unet_head import UNetHead2D
class UNet(AutoCfdModel):
def __init__(
self,
in_chan: int,
out_chan: int,
loss_fn: nn.Module,
n_case_params: int,
insert_case_params_at: str = "hidden",
bilinear: bool = False,
dim: int = 8,
):
assert insert_case_params_at in ["hidden", "input"]
super().__init__(loss_fn)
self.in_chan = in_chan
self.out_chan = out_chan
self.n_case_params = n_case_params
self.insert_case_params_at = insert_case_params_at
self.dim = dim
# 计算 Encoder 输入通道
encoder_in_chan = in_chan + 1 # + Mask
if insert_case_params_at == "input":
encoder_in_chan += n_case_params
# 1. Encoder
self.encoder = UNetEncoder2D(
in_channels=encoder_in_chan,
base_channels=dim,
num_stages=4,
bilinear=bilinear,
normtype="bn"
)
# 2. Hidden Injection
self.case_params_fc = None
if insert_case_params_at == "hidden":
bottleneck_dim = dim * 16
self.case_params_fc = nn.Linear(n_case_params, bottleneck_dim)
# 3. Decoder
self.decoder = UNetDecoder2D(
base_channels=dim,
num_stages=4,
bilinear=bilinear,
normtype="bn"
)
# 4. Head
self.head = UNetHead2D(
in_channels=dim,
out_channels=out_chan
)
def forward(
self,
inputs: Tensor,
case_params: Tensor,
mask: Optional[Tensor] = None,
label: Optional[Tensor] = None,
):
batch_size, n_chan, height, width = inputs.shape
residual = inputs[:, : self.out_chan]
# 构造 Mask
if mask is None:
mask = torch.ones((batch_size, 1, height, width)).to(inputs.device)
else:
if mask.dim() == 3:
mask = mask.unsqueeze(1) # (B, H, W) -> (B, 1, H, W)
# 拼接 Mask
x_in = torch.cat([inputs, mask], dim=1)
# 拼接 Case Params
if self.insert_case_params_at == "input":
cp_spatial = case_params.view(batch_size, self.n_case_params, 1, 1)
cp_spatial = cp_spatial.expand(-1, -1, height, width)
x_in = torch.cat([x_in, cp_spatial], dim=1)
# Encoder
features = self.encoder(x_in)
# Hidden 注入
if self.insert_case_params_at == "hidden":
bottleneck = features[-1]
conds = self.case_params_fc(case_params)
conds = conds.view(batch_size, -1, 1, 1)
features[-1] = bottleneck + conds
# Decoder
decoded = self.decoder(features)
# Head
preds = self.head(decoded)
# Residual & Mask
preds = preds + residual
preds = preds * mask
if label is not None:
label = label * mask
loss = self.loss_fn(labels=label, preds=preds)
return {"preds": preds, "loss": loss}
return {"preds": preds}
def generate_many(
self, inputs: Tensor, case_params: Tensor, mask: Tensor, steps: int
) -> List[Tensor]:
preds = []
# 处理单样本输入 (增加 Batch 维)
if inputs.dim() == 3:
inputs = inputs.unsqueeze(0)
case_params = case_params.unsqueeze(0)
if mask.dim() == 2:
mask = mask.unsqueeze(0)
# 确保 Mask 是 (B, 1, H, W) 以匹配 forward 逻辑
if mask.dim() == 3:
mask = mask.unsqueeze(1)
cur_frame = inputs
for _ in range(steps):
out_dict = self.forward(cur_frame, case_params=case_params, mask=mask)
cur_frame = out_dict["preds"]
preds.append(cur_frame)
return preds
def generate(
self,
inputs: Tensor,
case_params: Tensor,
mask: Optional[Tensor] = None,
) -> Tensor:
outputs = self.forward(inputs, case_params=case_params, mask=mask)
return outputs["preds"]
|