Faire des prédictions avec la multiplication de matrices
Dans les chapitres suivants, vous apprendrez à entraîner des modèles de régression linéaire. Ce processus produit un vecteur de paramètres qui peut être multiplié par les données d'entrée pour générer des prédictions. Dans cet exercice, vous utiliserez les données d'entrée, features, et un vecteur cible, bill, provenant d'un jeu de données de cartes de crédit que nous utiliserons plus tard dans le cours.
\(features = \begin{bmatrix} 2 & 24 \\ 2 & 26 \\ 2 & 57 \\ 1 & 37 \end{bmatrix}\), \(bill = \begin{bmatrix} 3913 \\ 2682 \\ 8617 \\ 64400 \end{bmatrix}\), \(params = \begin{bmatrix} 1000 \\ 150 \end{bmatrix}\)
La matrice de données d'entrée, features, contient deux colonnes : le niveau de scolarité et l'âge. Le vecteur cible, bill, correspond au montant de l'état de compte de l'emprunteur par carte de crédit.
Comme nous n'avons pas entraîné le modèle, vous allez proposer une valeur approximative pour le vecteur de paramètres, params. Vous utiliserez ensuite matmul() pour effectuer la multiplication de matrices de features par params afin de générer des prédictions, billpred, que vous comparerez à bill. Notez que nous avons importé matmul() et constant().
Cette activité fait partie du cours
Introduction à TensorFlow en Python
Instructions de l’exercice
- Définissez
features,paramsetbillcomme constantes. - Calculez le vecteur des valeurs prédites,
billpred, en multipliant les données d'entrée,features, par les paramètres,params. Utilisez la multiplication de matrices, et non le produit élément par élément. - Définissez
errorcomme les cibles,bill, moins les valeurs prédites,billpred.
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
# Define features, params, and bill as constants
features = ____([[2, 24], [2, 26], [2, 57], [1, 37]])
params = ____([[1000], [150]])
bill = ____([[3913], [2682], [8617], [64400]])
# Compute billpred using features and params
billpred = ____
# Compute and print the error
error = ____ - ____
print(error.numpy())