Showing posts with label python. Show all posts
Showing posts with label python. Show all posts

Python Seaborn - Matrix Plots

 # Seaborn — Matrix Plots

# Matrix plot = a grid where rows and columns are categories/variables
# and each cell is colored based on a value
# Used to see patterns, correlations, and relationships at a glance
# Two main matrix plots: heatmap and clustermap
import seaborn as sns
import matplotlib.pyplot as plt
import pandas as pd
import numpy as np
# ══════════════════════════════════════════════════════════════════════════════
# ── 1. heatmap — Color Grid showing values ────────────────────────────────────
# Each cell = a number shown as a color
# Darker/brighter color = higher or lower value (depends on colormap)
# Great for: showing correlation between columns, confusion matrix, pivot tables
# ══════════════════════════════════════════════════════════════════════════════

# ── Simple heatmap from a 2D list ─────────────────────────────────────────────
data = pd.DataFrame(
[[10, 20, 30],
[40, 50, 60],
[70, 80, 90]],
index=["Row A", "Row B", "Row C"], # row labels (y-axis)
columns=["Col 1", "Col 2", "Col 3"] # column labels (x-axis)
)
# data=pd.DataFrame({
# "col1" : [10,40,70],
# "col2" : [20,50,80],
# "col3" : [30,60,90]
# })
# data.index=["Row A","Row B","Row C"]
sns.heatmap(data)
plt.title("Basic Heatmap")
plt.show() # Output: 3×3 color grid — darker cells = higher values (bottom row darkest)

# ── annot=True — show numbers inside each cell ────────────────────────────────
sns.heatmap(data, annot=True)
plt.title("Heatmap with Numbers")
plt.show() # Output: same color grid but each cell also shows its number

# ── fmt= — format of numbers shown inside cells ───────────────────────────────
sns.heatmap(data, annot=True, fmt="d") # d = integer format (no decimals)
plt.title("Heatmap with Integer Labels")
plt.show() # Output: numbers shown as 10, 20, 30 (not 10.0, 20.0...)

sns.heatmap(data, annot=True, cmap="Blues") # light blue → dark blue
plt.title("Heatmap with color theme Blues")
plt.show() # Output: low values = light blue, high values = dark blue

sns.heatmap(data, annot=True, cmap="YlOrRd") # yellow → orange → red
plt.title("Heatmap with color theme YlOrRd")
plt.show() # Output: low values = yellow, high values = red

sns.heatmap(data, annot=True, cmap="coolwarm") # blue → white → red
plt.title("Heatmap with color theme coolwarm")
plt.show() # Output: low = blue, middle = white, high = red (good for correlation)

sns.heatmap(data, annot=True, cmap="Greens") # light green → dark green
plt.title("Heatmap with color theme Greens")
plt.show() # Output: low values = light green, high values = dark green

# ── linewidths= — adds borders between cells ──────────────────────────────────
sns.heatmap(data, annot=True, linewidths=0.5, linecolor="white")
plt.title("Heatmap with Cell Borders")
plt.show() # Output: white lines separate each cell — easier to read

# ── vmin / vmax — fix the color scale range ───────────────────────────────────
# vmin = value mapped to the lightest color
# vmax = value mapped to the darkest color
sns.heatmap(data, annot=True, vmin=0, vmax=100)
plt.title("Heatmap with Fixed Color Scale (0 to 100)")
plt.show() # Output: colors scale from 0 (light) to 100 (dark) — 90 is near darkest
# ══════════════════════════════════════════════════════════════════════════════
# ── 2. Correlation Heatmap — most common real-world use ───────────────────────
# Correlation = how much two columns move together
# +1.0 → both increase together (perfect positive)
# 0.0 → no relationship
# -1.0 → one increases while other decreases (perfect negative)
# df.corr() calculates correlation between all numeric columns
# ══════════════════════════════════════════════════════════════════════════════
df = pd.DataFrame({
"age": [22, 25, 30, 35, 40, 45, 50],
"salary": [30000, 35000, 50000, 60000, 72000, 80000, 90000],
"experience": [1, 2, 5, 8, 12, 18, 25],
"score": [85, 80, 75, 70, 65, 60, 55]
})

corr = df.corr() # corr() returns a table of correlation values between every column pair
#calculates how strongly each pair of columns is related to each other.

sns.heatmap(corr,
annot=True, # show correlation values inside cells
fmt=".2f", # 2 decimal places e.g. 0.98
cmap="coolwarm", # blue=negative, red=positive correlation
vmin=-1, vmax=1) # fix scale from -1 to +1

plt.title("Correlation Heatmap")
plt.show() # Output: grid showing how strongly each pair of columns is related
# age vs salary = ~0.99 (strong positive), age vs score = ~-0.99 (strong negative)

# ── mask= — hide the upper triangle (avoid duplicate info) ────────────────────
# Correlation table is symmetric — top-right mirrors bottom-left
# mask hides the duplicate upper triangle so it's easier to read
mask = np.zeros_like(corr, dtype=bool) # start with all False (show everything)
mask[np.triu_indices_from(mask)] = True # set upper triangle to True (hide it)

sns.heatmap(corr, annot=True, fmt=".2f", cmap="coolwarm", mask=mask, vmin=-1, vmax=1)
plt.title("Correlation Heatmap — Lower Triangle Only")
plt.show() # Output: only bottom-left half shown — cleaner, no repeated values
# ══════════════════════════════════════════════════════════════════════════════
# ── 3. clustermap — Heatmap with automatic grouping (clustering) ──────────────
# Same as heatmap BUT it reorders rows and columns automatically
# so that similar rows/columns are placed next to each other
# Dendrograms (tree diagrams) on top and left show which rows/cols are similar
# ══════════════════════════════════════════════════════════════════════════════
# ── Simple clustermap ────────────────────────────────────────────────────────
data2 = pd.DataFrame({
"Math": [90, 85, 40, 45, 70],
"Science": [88, 80, 42, 50, 68],
"History": [45, 50, 85, 90, 55],
"Art": [40, 45, 88, 92, 60]
}, index=["Alice", "Bob", "Carol", "Dave", "Eve"])

sns.clustermap(data2)
plt.suptitle("Clustermap — similar students and subjects grouped together", y=1.02)
plt.show() # Output: heatmap with rows/cols reordered — students good at Math/Science
# grouped together, students good at History/Art grouped together
# ── annot=True — show values in cells ────────────────────────────────────────
sns.clustermap(data2, annot=True, fmt="d", cmap="YlOrRd")
plt.suptitle("Clustermap with Values", y=1.02)
plt.show() # Output: colored grid with scores shown, similar rows/cols clustered
# ── standard_scale= — normalize data before clustering ───────────────────────
# standard_scale=1 → scale each column so values go from 0 to 1
# Useful when columns have very different ranges (e.g. salary vs age)
sns.clustermap(data2, standard_scale=1, cmap="Blues")
plt.suptitle("Clustermap with Normalized Columns (0 to 1)", y=1.02)
plt.show() # Output: each column scaled 0–1, makes comparison fair across columns
# ── z_score= — normalize by row or column using z-score ──────────────────────
# z_score=0 → normalize each row z_score=1 → normalize each column
# Shows which values are above/below average within each row or column
sns.clustermap(data2, z_score=1, cmap="coolwarm")
plt.suptitle("Clustermap with Z-score (above/below average per column)", y=1.02)
plt.show() # Output: blue = below average, red = above average within each subject
# ══════════════════════════════════════════════════════════════════════════════
# ── heatmap vs clustermap ─────────────────────────────────────────────────────
# ┌─────────────┬──────────────────────────────────┬────────────────────────────────┐
# │ │ heatmap │ clustermap │
# ├─────────────┼──────────────────────────────────┼────────────────────────────────┤
# │ Row order │ stays as-is │ reordered to group similar rows│
# │ Col order │ stays as-is │ reordered to group similar cols│
# │ Dendrogram │ no │ yes (tree on top and left) │
# │ Best for │ fixed grids like confusion matrix│ finding hidden patterns/groups │
# └─────────────┴──────────────────────────────────┴────────────────────────────────┘
# ══════════════════════════════════════════════════════════════════════════════

# ══════════════════════════════════════════════════════════════════════════════
# ── Quick Reference ───────────────────────────────────────────────────────────
# ══════════════════════════════════════════════════════════════════════════════

# sns.heatmap(data) → color grid from a 2D table
# sns.heatmap(data, annot=True) → show values inside each cell
# sns.heatmap(data, fmt="d") → integer format inside cells
# sns.heatmap(data, fmt=".2f") → 2 decimal format inside cells
# sns.heatmap(data, cmap="coolwarm") → set color theme
# sns.heatmap(data, vmin=0, vmax=1) → fix color scale range
# sns.heatmap(data, linewidths=0.5) → borders between cells
# sns.heatmap(data, mask=mask) → hide certain cells (e.g. upper triangle)
# df.corr() → correlation table between all numeric columns
#
# sns.clustermap(data) → heatmap with auto-grouping of similar rows/cols
# sns.clustermap(data, standard_scale=1) → normalize columns 0 to 1
# sns.clustermap(data, z_score=1) → show above/below average per column
#
# plt.show() → display the plot

Python Seaborn - Distribution Plots

 # Seaborn — a Python library built on top of Matplotlib

# Makes statistical plots easier and prettier with less code
# import convention: import seaborn as sns


# Distribution Plot — shows how data is spread / how often values occur
# Useful to understand: shape, center, spread, and outliers of data

import seaborn as sns
import matplotlib.pyplot as plt
import numpy as np

# ── Sample data used throughout ───────────────────────────────────────────────
ages = [22, 25, 25, 27, 28, 30, 30, 30, 32, 35, 35, 38, 40, 42, 45]
scores = np.random.seed(42) or np.random.normal(loc=70, scale=10, size=200)
# loc=70 → mean is 70, scale=10 → std deviation 10, size=200 → 200 values

# ══════════════════════════════════════════════════════════════════════════════
# ── 1. histplot — Histogram ───────────────────────────────────────────────────
# Shows how many times each value (or range of values) appears
# x-axis: value ranges (bins), y-axis: count of values in that range
# ══════════════════════════════════════════════════════════════════════════════

sns.histplot(ages, bins=5, color="steelblue") # bins=5 → divide data into 5 intervals
plt.title("Age Distribution")
plt.xlabel("Age")
plt.ylabel("Count")
plt.show()
# ── bins — controls the number of intervals ───────────────────────────────────
sns.histplot(ages, bins=3) # fewer bins → wider bars, less detail
plt.show() # Output: 3 tall wide bars
sns.histplot(ages, bins=10) # more bins → narrower bars, more detail
plt.show() # Output: 10 narrow bars showing finer breakdown

# ── kde=True — adds a smooth curve over the histogram ────────────────────────
# KDE = Kernel Density Estimate — a smooth line showing the shape of distribution
sns.histplot(ages, bins=5, kde=True, color="teal")
plt.title("Histogram with KDE Curve")
plt.show() # Output: bars + a smooth curved line on top

# ── stat= — changes what y-axis shows ────────────────────────────────────────
agess = [10,20,30,10,20,30,40]
plt.title("changes what y-axis shows")
sns.histplot(agess, stat="count") # count → number of values (default)
sns.histplot(agess, stat="frequency") # frequency → proportion per bin width
sns.histplot(agess, stat="density") # density → area under curve = 1 (for KDE)
sns.histplot(agess, stat="probability") # probability → each bar = fraction of total
plt.show() # Output: y-axis changes based on stat used


# ══════════════════════════════════════════════════════════════════════════════
# ── 2. kdeplot — Smooth Density Curve ────────────────────────────────────────
# KDE = Kernel Density Estimate
# Instead of bars, shows a smooth curve — great for seeing the shape of data
# y-axis shows density (not count) — area under the curve = 1
# ══════════════════════════════════════════════════════════════════════════════

np.random.seed(42)
# np.random.normal(loc, scale, size) — generates random numbers that cluster around a center
# normal is used here specifically because KDE plots are meant to show bell-shaped distributions.
# loc=70 → center / average — most numbers will be near 70
# scale=10 → spread — how far numbers go from center
# ~68% of values fall between 60–80 (70 ± 10)
# ~95% of values fall between 50–90 (70 ± 20)
# size=200 → generate 200 numbers
# Think of it as: exam scores for 200 students — most score around 70, few very high or low
scores = np.random.normal(loc=70, scale=10, size=200)
sns.kdeplot(scores, color="blue")
plt.title("Score Distribution (KDE)")
plt.xlabel("Score")
plt.ylabel("Density")
plt.show()

# ── fill=True — fills area under the curve ───────────────────────────────────
sns.kdeplot(scores, fill=True, color="skyblue", alpha=0.6) # alpha = transparency
plt.title("KDE with Fill")
plt.show() # Output: filled blue area under the curve
# ── bw_adjust — controls smoothness of the curve ─────────────────────────────
# bw_adjust < 1 → more detail / jagged, bw_adjust > 1 → smoother / wider
sns.kdeplot(scores, bw_adjust=0.5, label="Less smooth (0.5)")
sns.kdeplot(scores, bw_adjust=2.0, label="More smooth (2.0)")
plt.legend()
plt.title("KDE Smoothness Comparison")
plt.show() # Output: two curves — one tighter, one wider/smoother

# ── Multiple KDE curves on same plot ─────────────────────────────────────────
group_a = np.random.normal(60, 8, 100) # Group A: mean=60
group_b = np.random.normal(80, 10, 100) # Group B: mean=80

sns.kdeplot(group_a, label="Group A", fill=True, alpha=0.4)
sns.kdeplot(group_b, label="Group B", fill=True, alpha=0.4)
plt.legend()
plt.title("Two Groups Compared")
plt.show() # Output: two overlapping filled curves — easy to compare groups

# distplot(data) → sns.histplot(data, kde=True) ✔
# distplot(data, hist=False)→ sns.kdeplot(data) ✔
# distplot(data, kde=False) → sns.histplot(data) ✔
np.random.seed(42)
scores = np.random.normal(loc=70, scale=10, size=200)

sns.histplot(scores, kde=True, color="steelblue")
plt.title("histplot + kde=True")
plt.show() # Output: bars with smooth curve on top — same look as old distplot

sns.kdeplot(scores, fill=True, color="teal")
plt.title("kdeplot (replaces distplot kde-only mode)")
plt.show() # Output: filled smooth density curve

# ══════════════════════════════════════════════════════════════════════════════
# ── 4. displot — Flexible Distribution Plot (combines hist + kde) ─────────────
# displot = distribution plot — one function that can draw hist, kde, or ecdf
# kind= switches between types
# ══════════════════════════════════════════════════════════════════════════════
np.random.seed(42)
scores = np.random.normal(loc=70, scale=10, size=200)
sns.displot(scores, kind="hist") # same as histplot
plt.show() # Output: histogram bars

sns.displot(scores, kind="kde") # same as kdeplot
plt.show() # Output: smooth density curve

sns.displot(scores, kind="ecdf") # ECDF = shows what % of data is below each value
plt.show() # Output: S-shaped step curve going from 0% to 100%distplot
# ── hist + kde together ───────────────────────────────────────────────────────
sns.displot(scores, kind="hist", kde=True, color="mediumseagreen")
plt.title("Histogram + KDE using displot")
plt.show() # Output: bars with smooth line on top

# ══════════════════════════════════════════════════════════════════════════════
# ── 5. ecdfplot — Cumulative Distribution ────────────────────────────────────
# ECDF = Empirical Cumulative Distribution Function
# For each value on x-axis, shows: what fraction of data is ≤ that value
# Useful to answer: "what % of students scored below 75?"
# ══════════════════════════════════════════════════════════════════════════════
np.random.seed(42)
scores = np.random.normal(loc=70, scale=10, size=200)

sns.ecdfplot(scores, color="darkorange")
plt.title("Cumulative Score Distribution")
plt.xlabel("Score")
plt.ylabel("Proportion (0 to 1)")
plt.axhline(0.5, color="gray", linestyle="--") # horizontal line at 50%
plt.show() # Output: S-curve; where it crosses 0.5 = median score
# ══════════════════════════════════════════════════════════════════════════════
# ── 6. rugplot — Data Point Markers ──────────────────────────────────────────
# Draws a small tick mark on the x-axis for every single data point
# Shows exactly where each value sits — useful combined with kde or hist
# ══════════════════════════════════════════════════════════════════════════════

small_data = [22, 25, 28, 30, 30, 33, 35, 38, 40]

sns.kdeplot(small_data, color="blue")
sns.rugplot(small_data, color="red", height=0.05) # height = tick mark size
plt.title("KDE + Rug Plot")
plt.xlabel("Value")
plt.show() # Output: smooth curve with tiny red ticks below it showing raw values

# ══════════════════════════════════════════════════════════════════════════════
# ── 7. jointplot — Relationship + Distribution of Two Variables ───────────────
# Shows center plot (relationship between x and y)
# + side plots (distribution of x alone and y alone)
# Useful to see: how two variables relate AND how each is spread
# ══════════════════════════════════════════════════════════════════════════════

import pandas as pd

np.random.seed(42)
age = np.random.randint(22, 55, 100) # 100 random ages between 22 and 55
salary = age * 1500 # salary = age × 1500 (older = higher salary)

df = pd.DataFrame({"age": age, "salary": salary})

# ── kind="scatter" — dots in center (default) ────────────────────────────────
sns.jointplot(x="age", y="salary", data=df, kind="scatter")
plt.suptitle("scatter — dots showing each person", y=1.02)
plt.show() # Output: scatter dots in center, histogram of age on top, salary on right

# ── kind="kde" — smooth density contours ─────────────────────────────────────
# Contour lines show where most data points are concentrated (like a topographic map)
sns.jointplot(x="age", y="salary", data=df, kind="kde")
plt.suptitle("kde — density contours (where most points are)", y=1.02)
plt.show() # Output: oval contour rings in center, smooth kde curves on sides

# ── kind="hist" — 2D histogram grid ──────────────────────────────────────────
# Divides the space into a grid of squares — darker square = more points in that area
sns.jointplot(x="age", y="salary", data=df, kind="hist")
plt.suptitle("hist — grid squares, darker = more values", y=1.02)
plt.show() # Output: colored grid squares in center, histograms on sides

# ── kind="hex" — hexagonal bins ──────────────────────────────────────────────
# Like hist but uses hexagons — better when many points overlap
# Darker hexagon = more data points in that area
sns.jointplot(x="age", y="salary", data=df, kind="hex")
plt.suptitle("hex — hexagons, darker = more values (good for large data)", y=1.02)
plt.show() # Output: hexagonal grid in center, histograms on sides

# ── kind="reg" — scatter + regression line ───────────────────────────────────
# regression line = a straight line showing the overall trend in the data
sns.jointplot(x="age", y="salary", data=df, kind="reg")
plt.suptitle("reg — scatter with trend line", y=1.02)
plt.show() # Output: dots with a best-fit line, kde curves on sides



# ══════════════════════════════════════════════════════════════════════════════
# ── 8. pairplot — Distribution + Relationship for All Column Pairs ────────────
# Automatically plots every combination of columns in a DataFrame
# Diagonal → distribution of each column (how data is spread)
# Off-diagonal → relationship between each pair of columns (scatter)
#
# Example with 3 columns (age, salary, score):
# age salary score
# age [hist] [scatter] [scatter]
# salary[scatter] [hist] [scatter]
# score [scatter] [scatter] [hist]
# ══════════════════════════════════════════════════════════════════════════════


np.random.seed(42)
df2 = pd.DataFrame({
"age": np.random.randint(22, 55, 50), # 50 random ages
"salary": np.random.randint(22, 55, 50) * 1500, # salary based on age
"score": np.random.randint(50, 100, 50), # random exam scores
"dept": np.random.choice(["HR", "IT", "Finance"], 50) # random department
})

# ── Basic pairplot ────────────────────────────────────────────────────────────
sns.pairplot(df2[["age", "salary", "score"]]) # pass only numeric columns
plt.suptitle("Basic Pairplot", y=1.02)
plt.show() # Output: 3×3 grid — diagonal has histograms, others have scatter dots

# ── diag_kind="kde" — smooth curve on diagonal instead of histogram ───────────
sns.pairplot(df2[["age", "salary", "score"]], diag_kind="kde")
plt.suptitle("Pairplot with KDE on diagonal", y=1.02)
plt.show() # Output: 3×3 grid — diagonal has smooth curves, others have scatter dots

# ── kind="reg" — scatter + trend line on off-diagonal ────────────────────────
sns.pairplot(df2[["age", "salary", "score"]], kind="reg")
plt.suptitle("Pairplot with regression lines", y=1.02)
plt.show() # Output: each scatter plot has a best-fit trend line drawn through it

# ── hue= — color-code points by a category column ────────────────────────────
# hue splits the data by a category and colors each group differently
sns.pairplot(df2, hue="dept") # color by department
plt.suptitle("Pairplot colored by Department", y=1.02)
plt.show() # Output: dots colored by HR/IT/Finance — easy to compare groups


# ══════════════════════════════════════════════════════════════════════════════
# ── Quick Reference ───────────────────────────────────────────────────────────
# ══════════════════════════════════════════════════════════════════════════════

# sns.histplot(data) → histogram bars (count per range)
# sns.histplot(data, kde=True) → histogram + smooth curve on top
# sns.histplot(data, bins=n) → control number of intervals
# sns.histplot(data, stat="density") → y-axis as density instead of count
#
# sns.kdeplot(data) → smooth density curve only
# sns.kdeplot(data, fill=True) → fill area under curve
# sns.kdeplot(data, bw_adjust=0.5) → adjust smoothness
#
# sns.distplot(data) → DEPRECATED ❌ — use histplot/kdeplot instead
# sns.histplot(data, kde=True) → modern replacement for distplot
#
# sns.displot(data, kind="hist") → histogram via displot
# sns.displot(data, kind="kde") → kde via displot
# sns.displot(data, kind="ecdf") → cumulative curve via displot
#
# sns.ecdfplot(data) → cumulative distribution (0 to 1)
# sns.rugplot(data) → tick marks for each data point
#
# sns.jointplot(x=, y=, data=) → relationship + distribution of 2 variables
# sns.jointplot(..., kind="scatter") → dots in center (default)
# sns.jointplot(..., kind="kde") → smooth density contours
# sns.jointplot(..., kind="hist") → 2D histogram grid
# sns.jointplot(..., kind="hex") → hexagonal bins (good for many overlapping points)
# sns.jointplot(..., kind="reg") → scatter + regression line

Python seaborn - Regression Plots

# Seaborn — Regression Plots
# Regression plot = scatter plot with a trend line drawn through the data
# Trend line shows the overall direction/pattern in the data
# Used to answer: "as X increases, what happens to Y?"
#
# Two main functions:
# regplot → single plot, simple to use
# lmplot → grid support, can split by category using col= row= hue=

import seaborn as sns
import matplotlib.pyplot as plt
import pandas as pd
# ── Sample data used throughout ───────────────────────────────────────────────
df = pd.DataFrame({
"age": [22, 25, 28, 30, 33, 35, 38, 40, 43, 45],
"salary": [30000, 35000, 42000, 50000, 55000, 62000, 70000, 75000, 82000, 90000],
"score": [90, 85, 80, 78, 74, 70, 65, 60, 55, 50],
"dept": ["HR","HR","IT","IT","HR","Finance","Finance","IT","HR","Finance"]
})

# ══════════════════════════════════════════════════════════════════════════════
# ── 1. regplot — Scatter + Trend Line ────────────────────────────────────────
# regplot draws:
# dots → each data point
# line → best fit trend line through the data
# shaded band around line → confidence interval (how reliable the line is)
# wider band = less confident, narrower = more confident
# ══════════════════════════════════════════════════════════════════════════════

sns.regplot(x="age", y="salary", data=df)
plt.title("Age vs Salary with Trend Line")
plt.show() # Output: dots going up-right with a rising trend line — salary increases with age

# ── ci= — confidence interval band ───────────────────────────────────────────
# ci = confidence interval — the shaded area around the trend line
# ci=95 (default) → 95% confident the true line is within this band
# ci=None → hide the shaded band
sns.regplot(x="age", y="salary", data=df, ci=95)
plt.title("regplot with 95% Confidence Band")
plt.show() # Output: trend line with shaded band — shows uncertainty of the line

sns.regplot(x="age", y="salary", data=df, ci=None)
plt.title("regplot without Confidence Band")
plt.show() # Output: clean trend line and dots, no shaded area

# ── color and marker styling ──────────────────────────────────────────────────
sns.regplot(x="age", y="salary", data=df,
color="green", # color of dots and line
scatter_kws={"color": "blue", "s": 80}, # scatter_kws — style the dots separately
line_kws={"color": "red", "linewidth": 2}) # line_kws — style the line separately
plt.title("Styled regplot")
plt.show() # Output: blue dots, red trend line

# ── negative trend — score decreases as age increases ────────────────────────
sns.regplot(x="age", y="score", data=df)
plt.title("Age vs Score — Negative Trend")
plt.show() # Output: dots going down-right with a falling trend line — score drops with age

# ══════════════════════════════════════════════════════════════════════════════
# ── 2. lmplot — regplot with grid support ────────────────────────────────────
# lm = linear model
# lmplot works exactly like regplot BUT supports:
# hue= → separate trend line per category (different colors)
# col= → one plot per category, side by side
# row= → one plot per category, stacked
# ══════════════════════════════════════════════════════════════════════════════

# ── basic lmplot — same output as regplot ─────────────────────────────────────
sns.lmplot(x="age", y="salary", data=df)
plt.title("Basic lmplot")
plt.show() # Output: scatter + trend line — same as regplot

# ── hue= — separate trend line per department ────────────────────────────────
# draws one line per unique value in the hue column — each in a different color
sns.lmplot(x="age", y="salary", data=df, hue="dept")
plt.title("Trend Line per Department")
plt.show() # Output: dots and lines colored by dept — HR/IT/Finance each get their own line

# ── col= — one plot per department, side by side ─────────────────────────────
sns.lmplot(x="age", y="salary", data=df, col="dept")
plt.suptitle("Age vs Salary — one plot per Dept", y=1.02)
plt.show() # Output: 3 separate plots side by side — one trend line per dept

# ── col= + hue= — split by dept, color by dept ───────────────────────────────
sns.lmplot(x="age", y="salary", data=df, col="dept", hue="dept")
plt.suptitle("Dept split + colored", y=1.02)
plt.show() # Output: 3 plots side by side, each dept's line and dots in its own color

# ── height= and aspect= — control plot size ───────────────────────────────────
sns.lmplot(x="age", y="salary", data=df, col="dept", height=4, aspect=1)
plt.suptitle("Larger plots per Dept", y=1.02)
plt.show() # Output: same 3 plots but each is taller and wider
# ══════════════════════════════════════════════════════════════════════════════
# ── 3. residplot — Shows errors of the trend line ────────────────────────────
# residual = how far each actual data point is from the trend line
# dot above zero line → actual value is HIGHER than what line predicted
# dot below zero line → actual value is LOWER than what line predicted
# If dots are randomly scattered → trend line is a good fit
# If dots show a pattern → trend line is missing something
# ══════════════════════════════════════════════════════════════════════════════

sns.residplot(x="age", y="salary", data=df, color="purple")
plt.axhline(0, color="gray", linestyle="--") # axhline — draws a horizontal line at 0
plt.title("Residual Plot — How far each point is from the trend line")
plt.xlabel("Age")
plt.ylabel("Residual (actual − predicted)")
plt.show() # Output: dots scattered above and below zero line — random = good fit
# ══════════════════════════════════════════════════════════════════════════════
# ── regplot vs lmplot ─────────────────────────────────────────────────────────
# ┌──────────────┬────────────────────────────────┬──────────────────────────────┐
# │ │ regplot │ lmplot │
# ├──────────────┼────────────────────────────────┼──────────────────────────────┤
# │ Basic use │ single scatter + trend line │ same as regplot │
# │ hue= │ not supported │ separate line per category │
# │ col= / row= │ not supported │ grid of plots by category │
# │ Best for │ quick single plot │ comparing groups / categories│
# └──────────────┴────────────────────────────────┴──────────────────────────────┘
# ══════════════════════════════════════════════════════════════════════════════

# ══════════════════════════════════════════════════════════════════════════════
# ── Quick Reference ───────────────────────────────────────────────────────────
# ══════════════════════════════════════════════════════════════════════════════

# sns.regplot(x=, y=, data=) → scatter + single trend line
# sns.regplot(..., ci=95) → confidence band (default 95%)
# sns.regplot(..., ci=None) → hide confidence band
# sns.regplot(..., scatter_kws={}) → style the dots (color, size)
# sns.regplot(..., line_kws={}) → style the line (color, width)
#
# sns.lmplot(x=, y=, data=) → same as regplot
# sns.lmplot(..., hue="col") → separate trend line per category
# sns.lmplot(..., col="col") → one plot per category (side by side)
# sns.lmplot(..., row="col") → one plot per category (stacked)
# sns.lmplot(..., height=4, aspect=1) → control plot size
#
# sns.residplot(x=, y=, data=) → shows how far each point is from trend line
# plt.axhline(0, linestyle="--") → draws zero reference line
#
# plt.show() → display the plot