| from itertools import product |
| from typing import List, Optional |
|
|
| import torch |
| from torch import Tensor |
|
|
| from .ffn import Ffn |
| from .base_model import AutoCfdModel |
| from .act_fn import get_act_fn |
| from .loss import MseLoss |
|
|
|
|
| class AutoFfn(AutoCfdModel): |
| """ |
| Equivalent to autoregressive data-driven PINN. |
| """ |
|
|
| def __init__( |
| self, |
| input_field_dim: int, |
| num_case_params: int, |
| query_dim: int, |
| loss_fn: MseLoss, |
| num_label_samples: int = 1000, |
| depth: int = 8, |
| width: int = 100, |
| act_norm: bool = False, |
| act_name="relu", |
| ): |
| """ |
| 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.input_field_dim = input_field_dim |
| self.num_case_params = num_case_params |
| self.query_dim = query_dim |
| self.depth = depth |
| self.width = width |
| self.act_name = act_name |
| self.act_norm = act_norm |
| self.num_label_samples = num_label_samples |
|
|
| self.in_dim = input_field_dim + num_case_params + query_dim |
| act_fn = get_act_fn(act_name, act_norm) |
| self.widths = [self.in_dim] + [width] * depth + [1] |
| self.ffn = Ffn( |
| self.widths, |
| act_fn=act_fn, |
| act_on_output=False, |
| ) |
|
|
| 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. |
| |
| ### Parameters |
| - `inputs: Tensor` -- (b, c, h, w) |
| - `labels: Tensor` -- (b, c, h, w) |
| - `query_idxs: Tensor` -- (k, 2), k is the number of query points, |
| each is an (x, y) coordinate. |
| - `mask: Tensor` -- Not used. |
| |
| ### Function |
| Input: [b, branch_dim + trunk_dim] |
| Output: [b, 1] |
| """ |
| batch_size, _num_chan, height, width = inputs.shape |
|
|
| |
| inputs = inputs[:, 0] |
| |
| flat_inputs = inputs.view(batch_size, -1) |
| flat_inputs = torch.cat( |
| [flat_inputs, case_params], dim=1 |
| ) |
|
|
| if query_idxs is None: |
| query_idxs = torch.tensor( |
| list(product(range(height), range(width))), |
| dtype=torch.long, |
| device=flat_inputs.device, |
| ) |
|
|
| n_queries = query_idxs.shape[0] |
|
|
| |
| |
| |
| flat_inputs = flat_inputs.repeat(n_queries, 1) |
| batch_query_idxs = query_idxs.repeat(batch_size, 1) |
|
|
| |
| flat_inputs = torch.cat([flat_inputs, batch_query_idxs.float()], dim=1) |
|
|
| preds = self.ffn(flat_inputs) |
| preds = preds.view(batch_size, -1) |
|
|
| |
| residuals = inputs[:, 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) |
| 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 |
| preds = [] |
| for _ in range(steps): |
| |
| cur_frame = self.generate( |
| cur_frame, case_params=case_params, mask=mask |
| ) |
| preds.append(cur_frame) |
| return preds |
|
|