在Google Colab中定义函数时遭遇‘name 'train_data' is not defined’错误的技术咨询
Let's break down how to resolve that name 'train_data' is not defined error, plus other hidden issues in your code that are preventing things from working correctly. I'll walk you through each fix step by step:
1. Why the train_data Error Happens
This error typically pops up if you mistakenly pass an undefined variable when calling plot_predictions() (e.g., writing plot_predictions(train_data=train_data) instead of using your existing X_train variable). However, your function has default parameters set, so the real root causes are the other issues in your code that break execution before you even get to calling the plot function.
2. Key Fixes for Your Code
Issue 1: model is Undefined Before Calling predict()
You're trying to run y_pred = model.predict(X_test) before defining or training any TensorFlow model. That's going to throw a name 'model' is not defined error first. Let's add a simple linear regression model (since your data is linear) to fix this:
# Define, compile, and train a basic linear model model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=[1]) ]) model.compile(loss=tf.keras.losses.mae, optimizer=tf.keras.optimizers.SGD(), metrics=["mae"]) # Train the model (verbose=0 hides training logs for clarity) model.fit(X_train, y_train, epochs=100, verbose=0) # Now generate predictions y_pred = model.predict(X_test)
Issue 2: Typo in Function Parameter Name
You misspelled predictions as peredictions in your function definition. This means when you try to use predictions inside the function, Python can't find the variable:
# Wrong parameter name def plot_predictions(train_data=X_train, train_labels=y_train, test_data=X_test, test_labels=y_test, peredictions=y_pred): # Corrected version def plot_predictions(train_data=X_train, train_labels=y_train, test_data=X_test, test_labels=y_test, predictions=y_pred):
Issue 3: Incorrect plt.scatter() Parameter
The plt.scatter() function uses label (singular) for legend labels, not labels (plural). Using the wrong name will break your legend:
# Wrong plt.scatter(train_data, train_labels, c="b", labels="Training data") # Corrected (do this for all three scatter calls) plt.scatter(train_data, train_labels, c="b", label="Training data") plt.scatter(test_data, test_labels, c="g", label="Testing data") plt.scatter(test_data, predictions, c="r", label="Predictions")
Issue 4: Duplicate Import
You imported matplotlib.pyplot as plt twice — just remove one of the lines to clean up your code.
3. Full Corrected Code
Here's the complete, working version of your code with all fixes applied:
import tensorflow as tf import numpy as np import matplotlib.pyplot as plt from tensorflow.keras.utils import plot_model X = tf.range(-100, 100, 4) y = X + 10 # Split data into train/test sets X_train = X[:40] # First 40 samples for training y_train = y[:40] X_test = X[40:] # Last 10 samples for testing y_test = y[40:] # Define, compile, and train the model model = tf.keras.Sequential([ tf.keras.layers.Dense(1, input_shape=[1]) ]) model.compile(loss=tf.keras.losses.mae, optimizer=tf.keras.optimizers.SGD(), metrics=["mae"]) model.fit(X_train, y_train, epochs=100, verbose=0) # Generate predictions y_pred = model.predict(X_test) def plot_predictions(train_data=X_train, train_labels=y_train, test_data=X_test, test_labels=y_test, predictions=y_pred): ''' Plots training data, test data, and compares predictions to true labels. ''' plt.figure(figsize=(10, 7)) # Plot training data in blue plt.scatter(train_data, train_labels, c="b", label="Training data") # Plot test data in green plt.scatter(test_data, test_labels, c="g", label="Testing data") # Plot predictions in red plt.scatter(test_data, predictions, c="r", label="Predictions") # Show legend plt.legend(); # Call the plotting function plot_predictions()
4. Verifying the train_data Error is Fixed
With the corrected code, calling plot_predictions() uses the default train_data=X_train parameter you set, so you won't get the "undefined" error anymore. If you ever want to pass custom data to the function, just make sure the variable you're passing is already defined (e.g., plot_predictions(train_data=my_custom_train_data)).
内容的提问来源于stack exchange,提问作者Ali Bahadorani

