| from itertools import product |
| from typing import List, Optional |
|
|
| import torch |
| from torch import nn, Tensor |
|
|
| from .ffn import Ffn |
| from .base_model import AutoCfdModel |
| from .act_fn import get_act_fn |
| from .loss import MseLoss |
|
|
|
|
| class CnnBranch(nn.Module): |
| def __init__( |
| self, in_chan: int, kernel_size: int, padding: int, depth: int = 4 |
| ): |
| super().__init__() |
| self.in_chan = in_chan |
| self.in_conv = nn.Conv2d( |
| in_chan, 32, kernel_size=kernel_size, padding=padding |
| ) |
| self.out_conv = nn.Conv2d( |
| 32, 32, kernel_size=kernel_size, padding=padding |
| ) |
|
|
| blocks = [] |
| for i in range(depth): |
| blocks += [ |
| nn.Conv2d(32, 32, kernel_size, padding=padding), |
| nn.MaxPool2d(2), |
| nn.ReLU(), |
| ] |
| self.blocks = nn.Sequential(*blocks) |
|
|
| def forward(self, x: Tensor) -> Tensor: |
| x = self.in_conv(x) |
| x = self.blocks(x) |
| x = self.out_conv(x) |
| return x |
|
|
|
|
| class AutoDeepONetCnn(AutoCfdModel): |
| """ |
| Auto-regressive DeepONet for CFD. |
| |
| Our task is different from the one that the original DeepONet. In the |
| original DeepONet, the input function (input to the branch net) |
| is the initial condition (IC), but here, we have a fixed (zero) IC. |
| Instead, we have different boundary conditions (BCs), but we also |
| want the model to predict the next time step given the current time step. |
| Ideally, we should have two different branch nets, one accepting the |
| BCs, one accepting the current time step. |
| |
| Here, we assume that the current time step includes the information about |
| BCs (which are the values on the bounds), so we just feed |
| the current time step to one branch net. |
| """ |
|
|
| def __init__( |
| self, |
| in_chan: int, |
| query_dim: int, |
| loss_fn: MseLoss, |
| height: int = 100, |
| width: int = 100, |
| num_case_params: int = 5, |
| trunk_depth: int = 4, |
| act_name="relu", |
| act_norm: bool = False, |
| act_on_output: bool = False, |
| ): |
| """ |
| Args: |
| - branch_dim: int, the dimension of the branch net input. |
| - trunk_dim: int, the dimension of the trunk net input. |
| """ |
| super().__init__(loss_fn) |
| self.in_chan = in_chan |
| self.query_dim = query_dim |
| self.num_case_params = num_case_params |
| self.trunk_depth = trunk_depth |
| self.height = height |
| self.width = width |
| self.act_name = act_name |
| self.act_norm = act_norm |
| self.act_on_output = act_on_output |
|
|
| act_fn = get_act_fn(act_name, act_norm) |
| self.trunk_dims = [query_dim] + [100] * trunk_depth + [4 * 4 * 32] |
| |
| self.branch_net = CnnBranch( |
| in_chan + 1 + num_case_params, kernel_size=5, padding=2 |
| ) |
| self.trunk_net = Ffn( |
| self.trunk_dims, act_fn=act_fn, act_on_output=False |
| ) |
| self.out_ffn = Ffn( |
| [32 * 4 * 4] * 3 + [1], act_fn=act_fn, act_on_output=False |
| ) |
| self.bias = nn.Parameter(torch.zeros(1)) |
|
|
| def forward( |
| self, |
| inputs: Tensor, |
| case_params: Tensor, |
| label: Optional[Tensor] = None, |
| mask: Optional[Tensor] = None, |
| query_idxs: Optional[Tensor] = None, |
| ): |
| """ |
| Here, we just randomly sample some points, and use the label values on |
| those points as the label. |
| |
| x: (b, c, h, w) |
| labels: (b, c, h, w) |
| query_point: (k, 2), k is the number of query points, each is an (x, y) |
| coordinate. |
| |
| Goal: |
| Input: [b, branch_dim + trunk_dim] |
| Output: [b, 1] |
| """ |
|
|
| |
| if mask is not None: |
| if mask.dim() == 3: |
| mask = mask.unsqueeze(1) |
| inputs = torch.cat([inputs, mask], dim=1) |
|
|
| batch_size, num_chan, height, width = inputs.shape |
|
|
| |
| residuals = inputs[:, :] |
|
|
| |
| case_params = case_params.unsqueeze(-1).unsqueeze(-1) |
| |
| case_params = case_params.expand( |
| -1, -1, inputs.shape[-2], inputs.shape[-1] |
| ) |
| inputs = torch.cat( |
| [inputs, case_params], dim=1 |
| ) |
|
|
| x_branch = self.branch_net(inputs) |
| x_branch = x_branch.view(batch_size, -1) |
|
|
| if query_idxs is None: |
| query_idxs = torch.tensor( |
| list(product(range(height), range(width))), |
| dtype=torch.long, |
| device=x_branch.device, |
| ) |
|
|
| |
| x_trunk = (query_idxs.float() - 50) / 100 |
| x_trunk = self.trunk_net(x_trunk) |
| x_trunk = x_trunk.unsqueeze(0) |
| x_branch = x_branch.unsqueeze(1) |
| |
| preds = x_branch * x_trunk |
| preds = self.out_ffn(preds) |
| preds = preds.squeeze(-1) |
|
|
| |
| residuals = residuals[ |
| :, 0, query_idxs[:, 0], query_idxs[:, 1] |
| ] |
| |
| preds += residuals |
|
|
| if label is not None: |
| label = label[:, 0] |
| |
| |
| labels = label[:, query_idxs[:, 0], query_idxs[:, 1]] |
| loss = self.loss_fn(labels=labels, preds=preds) |
| return dict( |
| preds=preds, |
| loss=loss, |
| ) |
|
|
| preds = preds.view(-1, 1, height, width) |
| return dict(preds=preds) |
|
|
| def generate( |
| self, inputs: Tensor, case_params: Tensor, mask: Tensor |
| ) -> Tensor: |
| """ |
| x: (c, h, w) or (B, c, h, w) |
| |
| Returns: |
| (b, c, h, w) |
| """ |
| if inputs.dim() == 3: |
| inputs = inputs.unsqueeze(0) |
| case_params = case_params.unsqueeze(0) |
| mask = mask.unsqueeze(0) |
| batch_size, num_chan, height, width = inputs.shape |
| query_idxs = torch.tensor( |
| list(product(range(height), range(width))), |
| dtype=torch.long, |
| device=inputs.device, |
| ) |
| |
| |
| preds = self.forward( |
| inputs, query_idxs=query_idxs, case_params=case_params, mask=mask |
| )["preds"] |
| preds = preds.view(-1, 1, height, width) |
| return preds |
|
|
| def generate_many( |
| self, inputs: Tensor, case_params: Tensor, mask: Tensor, steps: int |
| ) -> List[Tensor]: |
| """ |
| x: (c, h, w) or (B, c, h, w) |
| mask: (h, w). 1 for interior, 0 for boundaries. |
| steps: int, number of steps to generate. |
| |
| Returns: |
| list of tensors, each of shape (b, c, h, w) |
| """ |
| if inputs.dim() == 3: |
| inputs = inputs.unsqueeze(0) |
| case_params = case_params.unsqueeze(0) |
| mask = mask.unsqueeze(0) |
| cur_frame = inputs |
| p = inputs[:, -1:] |
| preds = [] |
| for _ in range(steps): |
| |
| cur_frame = self.generate( |
| cur_frame, case_params=case_params, mask=mask |
| ) |
| preds.append(cur_frame) |
| cur_frame = torch.cat([cur_frame, p], dim=1) |
| return preds |
|
|