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