MCPcopy Create free account
hub / github.com/tensorpack/tensorpack / get_config

Function get_config

examples/basics/cifar-convnet.py:108–128  ·  view source on GitHub ↗
(cifar_classnum)

Source from the content-addressed store, hash-verified

106
107
108def get_config(cifar_classnum):
109 # prepare dataset
110 dataset_train = get_data('train', cifar_classnum)
111 dataset_test = get_data('test', cifar_classnum)
112
113 def lr_func(lr):
114 if lr < 3e-5:
115 raise StopTraining()
116 return lr * 0.31
117 return TrainConfig(
118 model=Model(cifar_classnum),
119 data=QueueInput(dataset_train),
120 callbacks=[
121 ModelSaver(),
122 InferenceRunner(dataset_test,
123 ScalarStats(['accuracy', 'cost'])),
124 StatMonitorParamSetter('learning_rate', 'validation_accuracy', lr_func,
125 threshold=0.001, last_k=10, reverse=True),
126 ],
127 max_epoch=150,
128 )
129
130
131if __name__ == '__main__':

Callers 1

cifar-convnet.pyFile · 0.70

Calls 8

TrainConfigClass · 0.85
QueueInputClass · 0.85
ModelSaverClass · 0.85
InferenceRunnerClass · 0.85
ScalarStatsClass · 0.85
get_dataFunction · 0.70
ModelClass · 0.70

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…