python import numpy as np from pylearn2.models import mlp from pylearn2.training_algorithms import sgd from pylearn2.termination_criteria import EpochCounter from pylearn2.datasets import mnist from pylearn2.train import Train from pylearn2.training_callbacks import MonitorBasedSaveBest train_set = mnist.MNIST(which_set='train', start=0, stop=50000) valid_set = mnist.MNIST(which_set='train', start=50000, stop=60000) test_set = mnist.MNIST(which_set='test') layers = [mlp.Sigmoid(layer_name='h1', dim=100, irange=0.1), mlp.Softmax(layer_name='y', n_classes=10, irange=0.1)] model = mlp.MLP(layers, nvis=784) algorithm = sgd.SGD(learning_rate=0.01, cost=mlp.Costs.MeanSquaredReconstructionError()) termination_criterion = EpochCounter(max_epochs=10) callbacks = [MonitorBasedSaveBest(channel_name='valid_y_misclass', save_path='best_model.pkl')] trainer = Train(model=model, dataset=train_set, algorithm=algorithm, extensions=callbacks, termination_criterion=termination_criterion) trainer.main_loop() test_error = trainer.evaluate(test_set) print("Test Error: {:.2f}%".format(test_error * 100))


上一篇:
下一篇:
切换中文