CFDBench / model /__init__.py
OneScience's picture
Upload folder using huggingface_hub
338c3e4 verified
Raw
History Blame Contribute Delete
6.35 kB
from typing import Tuple
import torch
from model.auto_deeponet import AutoDeepONet
from model.auto_deeponet_cnn import AutoDeepONetCnn
from model.auto_edeeponet import AutoEDeepONet
from model.auto_ffn import AutoFfn
from model.deeponet import DeepONet
from model.ffn import FfnModel
from model.fno.fno2d import Fno2d
from model.loss import loss_name_to_fn
from model.resnet import ResNet
from model.unet import UNet
AUTO_MODEL_NAMES = {
"auto_ffn",
"auto_deeponet",
"auto_edeeponet",
"auto_deeponet_cnn",
"resnet",
"unet",
"fno",
}
STATIC_MODEL_NAMES = {"ffn", "deeponet"}
ALL_MODEL_NAMES = AUTO_MODEL_NAMES | STATIC_MODEL_NAMES
def infer_task_type(model_name: str) -> str:
if model_name in AUTO_MODEL_NAMES:
return "auto"
if model_name in STATIC_MODEL_NAMES:
return "static"
raise ValueError(f"Unknown CFDBench model.name={model_name!r}. Available: {sorted(ALL_MODEL_NAMES)}")
def get_input_shapes(data_name: str, num_rows: int, num_cols: int) -> Tuple[int, int, int]:
if any(name in data_name for name in ("tube", "dam", "cylinder")):
n_rows = num_rows + 2
n_cols = num_cols + 1
elif "cavity" in data_name:
n_rows = num_rows
n_cols = num_cols
else:
raise ValueError(f"Unknown CFDBench problem in data_name: {data_name}")
if "cylinder" in data_name:
n_case_params = 8
elif any(name in data_name for name in ("cavity", "tube", "dam")):
n_case_params = 5
else:
raise ValueError(f"Unknown CFDBench parameter set in data_name: {data_name}")
return n_rows, n_cols, n_case_params
def build_model(cfg) -> torch.nn.Module:
model_cfg = cfg.model
data_cfg = cfg.datapipe.data
source_cfg = cfg.datapipe.source
train_cfg = cfg.training
loss_fn = loss_name_to_fn(train_cfg.loss_name)
n_rows, n_cols, n_case_params = get_input_shapes(
data_name=source_cfg.data_name,
num_rows=data_cfg.num_rows,
num_cols=data_cfg.num_cols,
)
name = model_cfg.name
if name == "ffn":
widths = [n_case_params + 3] + [model_cfg.ffn_width] * model_cfg.ffn_depth + [1]
return FfnModel(
widths=widths,
loss_fn=loss_fn,
act_name=model_cfg.act_fn,
act_norm=model_cfg.act_scale_invariant,
act_on_output=model_cfg.act_on_output,
num_label_samples=model_cfg.num_label_samples,
)
if name == "deeponet":
return DeepONet(
branch_dim=n_case_params,
trunk_dim=3,
loss_fn=loss_fn,
width=model_cfg.deeponet_width,
trunk_depth=model_cfg.trunk_depth,
branch_depth=model_cfg.branch_depth,
act_name=model_cfg.act_fn,
act_norm=model_cfg.act_scale_invariant,
act_on_output=model_cfg.act_on_output,
num_label_samples=model_cfg.num_label_samples,
)
if name == "auto_ffn":
return AutoFfn(
input_field_dim=n_rows * n_cols,
num_case_params=n_case_params,
query_dim=2,
loss_fn=loss_fn,
width=model_cfg.autoffn_width,
depth=model_cfg.autoffn_depth,
act_name=model_cfg.act_fn,
act_norm=model_cfg.act_scale_invariant,
num_label_samples=model_cfg.num_label_samples,
)
if name == "auto_deeponet":
return AutoDeepONet(
branch_dim=n_rows * n_cols + n_case_params,
trunk_dim=2,
loss_fn=loss_fn,
width=model_cfg.deeponet_width,
trunk_depth=model_cfg.trunk_depth,
branch_depth=model_cfg.branch_depth,
act_name=model_cfg.act_fn,
act_norm=model_cfg.act_scale_invariant,
act_on_output=model_cfg.act_on_output,
num_label_samples=model_cfg.num_label_samples,
)
if name == "auto_edeeponet":
return AutoEDeepONet(
dim_branch1=n_rows * n_cols,
dim_branch2=n_case_params,
trunk_dim=2,
loss_fn=loss_fn,
width=model_cfg.autoedeeponet_width,
trunk_depth=model_cfg.autoedeeponet_depth,
branch_depth=model_cfg.autoedeeponet_depth,
act_name=model_cfg.autoedeeponet_act_fn,
num_label_samples=model_cfg.num_label_samples,
)
if name == "auto_deeponet_cnn":
if n_rows // 16 != 4 or n_cols // 16 != 4:
raise ValueError(
"auto_deeponet_cnn expects the post-padding grid to reduce to 4x4 after four 2x pools. "
"For tube/dam/cylinder fake data, set datapipe.data.num_rows=64 and num_cols=64."
)
return AutoDeepONetCnn(
in_chan=model_cfg.in_chan,
height=n_rows,
width=n_cols,
num_case_params=n_case_params,
query_dim=2,
loss_fn=loss_fn,
trunk_depth=model_cfg.trunk_depth,
act_name=model_cfg.act_fn,
act_norm=model_cfg.act_scale_invariant,
act_on_output=model_cfg.act_on_output,
)
if name == "resnet":
return ResNet(
in_chan=model_cfg.in_chan,
out_chan=model_cfg.out_chan,
n_case_params=n_case_params,
loss_fn=loss_fn,
hidden_chan=model_cfg.resnet_hidden_chan,
num_blocks=model_cfg.resnet_depth,
kernel_size=model_cfg.resnet_kernel_size,
padding=model_cfg.resnet_padding,
)
if name == "unet":
return UNet(
in_chan=model_cfg.in_chan,
out_chan=model_cfg.out_chan,
loss_fn=loss_fn,
n_case_params=n_case_params,
insert_case_params_at=model_cfg.unet_insert_case_params_at,
dim=model_cfg.unet_dim,
)
if name == "fno":
return Fno2d(
in_chan=model_cfg.in_chan,
out_chan=model_cfg.out_chan,
n_case_params=n_case_params,
loss_fn=loss_fn,
num_layers=model_cfg.fno_depth,
hidden_dim=model_cfg.fno_hidden_dim,
modes1=model_cfg.fno_modes_x,
modes2=model_cfg.fno_modes_y,
)
raise ValueError(f"Invalid CFDBench model.name={name!r}. Available: {sorted(ALL_MODEL_NAMES)}")