(cifar_classnum)
| 106 | |
| 107 | |
| 108 | def 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 | |
| 131 | if __name__ == '__main__': |
no test coverage detected
searching dependent graphs…