By the end of this lesson, you will be able to create scatter plots in Python using Matplotlib to visualize relationships between two numerical variables.
What it is
A scatter plot is a type of data visualization that displays individual data points on a two-dimensional plane. Each point represents an observation with coordinates defined by two variables: one plotted on the x-axis and the other on the y-axis. This chart type is fundamental for identifying correlations, clusters, or outliers in datasets. Related terms include correlation, linear regression, and data distribution. In Python, the primary library for creating these plots ismatplotlib.pyplot.
Why it matters
- Visualizing Correlations: Quickly determine if two variables move together (positive correlation), inversely (negative correlation), or independently.
- Detecting Outliers: Identify anomalous data points that deviate significantly from the general pattern.
- Understanding Distributions: See how data is spread across ranges, which helps in choosing appropriate statistical models.
- Exploratory Data Analysis (EDA): Serve as a first step in analyzing new datasets before applying complex algorithms.
Syntax or steps
The core function for creating a scatter plot isplt.scatter(). The basic workflow involves importing the library, preparing your data arrays, calling the function, and displaying the plot.
1. Import matplotlib.pyplot.
2. Define your x and y data lists or arrays.
3. Call plt.scatter(x, y).
4. Add labels and titles for clarity.
5. Display the plot using plt.show().
Example
import matplotlib.pyplot as plt
# Sample data: Hours studied vs. Exam scores
hours_studied = [1, 2, 3, 4, 5, 6, 7, 8]
exam_scores = [50, 55, 60, 65, 70, 75, 80, 85]
# Create the scatter plot
plt.scatter(hours_studied, exam_scores, color='blue', marker='o')
# Add labels and title
plt.xlabel('Hours Studied')
plt.ylabel('Exam Score')
plt.title('Relationship Between Study Time and Scores')
# Display the plot
plt.show()
Part-by-part explanation:
import matplotlib.pyplot as plt: Loads the plotting module.hours_studiedandexam_scores: Lists containing paired data points.plt.scatter(...): Plots each pair.color='blue'sets the dot color, andmarker='o'defines the shape (circle).plt.xlabel()andplt.ylabel(): Label the axes for context.plt.show(): Renders the figure window.
Common mistakes
- Mismatched Array Lengths: If
xandyhave different lengths, Matplotlib will raise an error. Ensure both lists contain the same number of elements. - Forgetting
plt.show(): Without this call, the plot may not appear in some environments (like scripts run from the terminal). - Overplotting: With very large datasets, points overlap, making patterns hard to see. Use transparency (
alpha=0.5) or hexbin plots instead. - Lack of Labels: A plot without axis labels is ambiguous. Always define what the x and y axes represent.
When to use it
Scatter plots are best for continuous numerical data. Compare them with line charts, which are better for time-series data where order matters.| Chart Type | Best For | Data Relationship |
|---|---|---|
| Scatter Plot | Correlation between two variables | No inherent order required |
| Line Chart | Trends over time | Sequential/ordered data |
| Bar Chart | Categorical comparisons | Discrete groups |
Practice
Guided Exercise: Modify the example above to change the marker size to 100 and add a grid usingplt.grid(True). Observe how the larger markers affect readability.
Challenge: Generate random data using numpy.random.rand(50) for both x and y. Plot it and try to identify any visual clusters. Hint: You may need to import numpy.
Quick check
Question: What parameter would you add toplt.scatter() to make overlapping points easier to distinguish?
Answer: Use alpha=0.5 to set transparency, allowing you to see density variations.
Summary
Scatter plots are essential tools for exploring relationships between numerical variables in Python. By masteringplt.scatter() and proper labeling, you can effectively communicate data insights such as correlations and outliers.