Actualizări multiple ale ponderilor
Acum vei face mai multe actualizări pentru a îmbunătăți semnificativ ponderile modelului și vei vedea cum se îmbunătățesc predicțiile cu fiecare actualizare.
Pentru a păstra codul curat, există o funcție pre-încărcată get_slope() care primește input_data, target și weights ca argumente. Există, de asemenea, o funcție get_mse() care primește aceleași argumente. Variabilele input_data, target și weights au fost pre-încărcate.
Această rețea nu are straturi ascunse și trece direct de la intrare (cu 3 noduri) la un nod de ieșire. Reține că weights este un singur array.
Am pre-încărcat și matplotlib.pyplot, iar istoricul erorilor va fi reprezentat grafic după ce ai efectuat pașii de gradient descent.
Acest exercițiu face parte din cursul
Introducere în Deep Learning în Python
Instrucțiuni pentru exercițiu
- Folosind o buclă
forpentru a actualiza iterativ ponderile:- Calculează panta folosind funcția
get_slope(). - Actualizează ponderile folosind o rată de învățare de
0.01. - Calculează eroarea medie pătratică (
mse) cu ponderile actualizate folosind funcțiaget_mse(). - Adaugă
mselamse_hist.
- Calculează panta folosind funcția
- Apasă Trimite răspunsul pentru a vizualiza
mse_hist. Ce tendință observi?
Exercițiu interactiv practic
Încearcă acest exercițiu completând acest cod de exemplu.
n_updates = 20
mse_hist = []
# Iterate over the number of updates
for i in range(n_updates):
# Calculate the slope: slope
slope = ____(____, ____, ____)
# Update the weights: weights
weights = ____ - ____ * ____
# Calculate mse with new weights: mse
mse = ____(____, ____, ____)
# Append the mse to mse_hist
____
# Plot the mse history
plt.plot(mse_hist)
plt.xlabel('Iterations')
plt.ylabel('Mean Squared Error')
plt.show()