You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

在Google Colab中定义函数时遭遇‘name 'train_data' is not defined’错误的技术咨询

Fixing "name 'train_data' is not defined" Error & Other Issues in Your TensorFlow Plotting Function

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.04.30 17:07:41