INNER CODE UNIT · Python
build_network
nv-tlabs/ATISS · scene_synthesis/networks/__init__.py:63
def build_network(
input_dims,
n_classes,
config,
weight_file=None,
device="cpu"):
network_type = config["network"]["type"]
if network_type == "autoregressive_transformer":
train_on_batch = train_on_batch_simple_autoregressive
validate_on_batch = validate_on_batch_simple_autoregressive
network = AutoregressiveTransformer(
input_dims,
hidden2output_layer(config, n_classes),
get_feature_extractor(
config["feature_extractor"].get("name", "resnet18"),
freeze_bn=config["feature_extractor"].get("freeze_bn", True),
input_channels=config["feature_extractor"].get("input_channels", 1),