MCPcopy Create free account
hub / github.com/MarcCoru/locationencoder / get_neural_network

Function get_neural_network

locationencoder/locationencoder.py:44–62  ·  view source on GitHub ↗
(name, input_dim, hparams=None)

Source from the content-addressed store, hash-verified

42 raise ValueError(f"{name} not a known positional encoding.")
43
44def get_neural_network(name, input_dim, hparams=None):
45 if name == "linear":
46 return nn.Linear(input_dim, hparams['num_classes'])
47 elif name == "siren":
48 return NN.SirenNet(
49 dim_in=input_dim,
50 dim_hidden=hparams['dim_hidden'],
51 num_layers=hparams['num_layers'],
52 dim_out=hparams['num_classes'],
53 dropout=hparams['dropout'] if "dropout" in hparams.keys() else False
54 )
55 elif name == "fcnet":
56 return NN.FCNet(
57 num_inputs=input_dim,
58 num_classes=hparams['num_classes'],
59 dim_hidden=hparams['dim_hidden']
60 )
61 else:
62 raise ValueError(f"{name} not a known neural networks.")
63
64def get_param(hparams, key, default=False):
65 """

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected