(args, n_inputs, n_outputs)
| 381 | |
| 382 | |
| 383 | def get_model(args, n_inputs, n_outputs): |
| 384 | nl, nh = args.n_layers, args.n_hiddens |
| 385 | if args.is_cifar: |
| 386 | net = ResNet18(n_outputs, bias=args.bias) |
| 387 | else: |
| 388 | net = MLP([n_inputs] + [nh] * nl + [n_outputs]) |
| 389 | return net |
| 390 | |
| 391 | |
| 392 | def main(overwrite_args=None): |