ÎncepețiÎncepe gratuit

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

Vezi cursul

Instrucțiuni pentru exercițiu

  • Folosind o buclă for pentru 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ția get_mse().
    • Adaugă mse la mse_hist.
  • 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()
Editează și rulează codul