| from typing import Optional, Dict, List |
| from itertools import product |
|
|
| from torch import nn, Tensor |
| import torch |
|
|
| from .base_model import CfdModel |
| from .loss import MseLoss |
| from .act_fn import get_act_fn |
|
|
|
|
| class Ffn(nn.Module): |
| """ |
| A general fully connected multi-layer neural network. |
| """ |
|
|
| def __init__( |
| self, dims: list, act_fn: nn.Module, act_on_output: bool = False |
| ): |
| super().__init__() |
| self.dims = dims |
|
|
| layers = [] |
| for i in range(len(dims) - 2): |
| layers.append(nn.Linear(dims[i], dims[i + 1])) |
| layers.append(act_fn) |
| |
| layers.append(nn.Linear(dims[-2], dims[-1])) |
| if act_on_output: |
| layers.append(act_fn) |
| self.layers = nn.Sequential(*layers) |
|
|
| def forward(self, x: Tensor) -> Tensor: |
| x = self.layers(x) |
| return x |
|
|
|
|
| class FfnModel(CfdModel): |
| """ |
| Non-autoregressive FNN for data-driven CFD. |
| |
| Branch net accepts the boundary and physics properties as inputs. |
| Trunk net accepts the query location (t, x, y) as input. |
| """ |
|
|
| def __init__( |
| self, |
| loss_fn: MseLoss, |
| widths: List[int], |
| act_name: str = "relu", |
| act_norm: bool = True, |
| act_on_output: bool = False, |
| num_label_samples: int = 1000, |
| ): |
| """ |
| 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.loss_fn = loss_fn |
| self.widths = widths |
| self.act_name = act_name |
| self.act_norm = act_norm |
| self.act_on_output = act_on_output |
| self.num_label_samples = num_label_samples |
|
|
| act_fn = get_act_fn(act_name, act_norm) |
| self.ffn = Ffn( |
| self.widths, act_fn=act_fn, act_on_output=self.act_on_output |
| ) |
|
|
| def forward( |
| self, |
| case_params: Tensor, |
| t: Tensor, |
| label: Optional[Tensor] = None, |
| query_idxs: Optional[Tensor] = None, |
| ) -> Dict[str, Tensor]: |
| """ |
| A faster forward by using all the points in the frame (`label`) at |
| time step `t` as training examples. |
| |
| Args: |
| - x_branch: (b, branch_dim), input to the branch net. |
| - t: (b), input to the trunk net, a batch of t |
| - label: (b, w, h), the frame to be predicted. |
| - query_idxs: (b, k, 2), the query locations. |
| """ |
|
|
| batch_size, dim_in = case_params.shape |
|
|
| if query_idxs is None: |
| |
| |
| assert label is not None |
| height, width = label.shape[-2:] |
| query_idxs = torch.stack( |
| [ |
| torch.randint( |
| 0, |
| height, |
| (self.num_label_samples,), |
| device=label.device, |
| ), |
| torch.randint( |
| 0, |
| width, |
| (self.num_label_samples,), |
| device=label.device, |
| ), |
| ], |
| dim=-1, |
| ) |
|
|
| |
| coords = query_idxs.unsqueeze(0) |
| coords = coords.repeat(batch_size, 1, 1) |
| num_queries = coords.shape[1] |
| t = t.unsqueeze(-1) |
| t = t.repeat(1, num_queries, 1) |
| coords = torch.cat([coords, t], dim=-1) |
|
|
| |
| case_params = case_params.unsqueeze(1) |
| case_params = case_params.repeat(1, num_queries, 1) |
| inp = torch.cat([case_params, coords], dim=-1) |
| inp = inp.view(batch_size * num_queries, -1) |
| preds = self.ffn(inp) |
| preds = preds.view(batch_size, num_queries) |
|
|
| if label is not None: |
| |
| label = label[:, 0] |
| labels = label[:, query_idxs[:, 0], query_idxs[:, 1]] |
| assert ( |
| preds.shape == labels.shape |
| ), f"{preds.shape}, {labels.shape}" |
| loss = self.loss_fn(preds=preds, labels=labels) |
| return dict( |
| preds=preds, |
| loss=loss, |
| ) |
| return dict( |
| preds=preds, |
| ) |
|
|
| def generate_one( |
| self, case_params: Tensor, t: Tensor, height: int, width: int |
| ) -> Tensor: |
| """ |
| Generate one frame at time t. |
| |
| Args: |
| - x_branch: Tensor, (branch_dim) |
| - t: Tensor, (1,) |
| - height: int |
| - width: int |
| |
| Returns: |
| (b, c, h, w) |
| """ |
| if len(case_params.shape) == 1: |
| case_params = case_params.unsqueeze(0) |
| if len(t.shape) == 0: |
| t = t.unsqueeze(0).unsqueeze(0) |
| elif len(t.shape) == 1: |
| t = t.unsqueeze(0) |
|
|
| |
| query_idxs = torch.tensor( |
| list(product(range(height), range(width))), |
| |
| device=case_params.device, |
| ) |
|
|
| |
| |
| output = self.forward(case_params, t=t, query_idxs=query_idxs)["preds"] |
| output = output.view(-1, 1, height, width) |
| return output |
|
|