¿Necesitamos más datos?
¡Es hora de comprobar si el model del conjunto de datos de dígitos que creaste se beneficia de tener más ejemplos de entrenamiento!
Para mantener el código al mínimo, ya hay varias cosas inicializadas y listas para usar:
- El
modelque acabas de construir. X_train,y_train,X_testyy_test.- Los
initial_weightsde tu modelo, guardados tras usarmodel.get_weights(). - Una lista predefinida de tamaños de entrenamiento:
training_sizes. - Un callback de early stopping predefinido que monitoriza la pérdida:
early_stop. - Dos listas vacías para almacenar los resultados de la evaluación:
train_accsytest_accs.
Entrena tu modelo con los distintos tamaños de entrenamiento y evalúa los resultados en X_test.
Termina representando los resultados con plot_results().
¡El código completo de este ejercicio está en las diapositivas!
Este ejercicio forma parte del curso
Introducción al Deep Learning con Keras
Instrucciones del ejercicio
- Obtén una fracción de los datos de entrenamiento determinada por el
sizeque estemos evaluando actualmente en el bucle. - Establece los pesos del modelo a
initial_weightsconset_weights()y entrena tu modelo con la fracción de datos de entrenamiento usandoearly_stopcomo callback. - Evalúa y guarda la accuracy para la fracción de entrenamiento y para el conjunto de prueba.
- Llama a
plot_results()pasando las accuracies de entrenamiento y prueba para cada tamaño de entrenamiento.
ejercicio interactivo práctico
Prueba este ejercicio completando este código de ejemplo.
for size in training_sizes:
# Get a fraction of training data (we only care about the training data)
X_train_frac, y_train_frac = X_train[:size], y_train[:size]
# Reset the model to the initial weights and train it on the new training data fraction
model.set_weights(____)
model.fit(X_train_frac, y_train_frac, epochs = 50, callbacks = [early_stop])
# Evaluate and store both: the training data fraction and the complete test set results
train_accs.append(model.evaluate(____, ____)[1])
test_accs.append(model.evaluate(____, ____)[1])
# Plot train vs test accuracies
plot_results(____, ____)