import torch.nn as nn def build_mlps(c_in, mlp_channels=None, ret_before_act=False, without_norm=False): layers = [] num_layers = len(mlp_channels) for k in range(num_layers): if k + 1 == num_layers and ret_before_act: layers.append(nn.Linear(c_in, mlp_channels[k], bias=True)) else: if without_norm: layers.extend([nn.Linear(c_in, mlp_channels[k], bias=True), nn.ReLU()]) else: layers.extend([nn.Linear(c_in, mlp_channels[k], bias=False), nn.BatchNorm1d(mlp_channels[k]), nn.ReLU()]) c_in = mlp_channels[k] return nn.Sequential(*layers)