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),

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…