import numpy as np
import matplotlib.pyplot as plt
# Define a simple 2D example
# Hyperplane:
# w^T x + b = 0
# x1 + x2 - 1 = 0
w = np.array([1.0, 1.0])
b = -1.0
# Scale factor
a = 3.0
w_scaled = a * w
b_scaled = a * b
# Points and labels
# y = +1 means the point should be above the boundary
# y = -1 means the point should be below the boundary
X = np.array([
[1.4, 1.0], # correctly classified positive
[0.2, 0.2], # correctly classified negative
[0.6, 0.4], # exactly on boundary
[0.3, 1.2], # misclassified negative
[1.1, 0.1], # correctly classified positive
])
y = np.array([1, -1, 1, -1, 1])
point_names = ["A", "B", "C", "D", "E"]
def decision_score(X, w, b):
"""
Compute f(x) = w^T x + b.
"""
return X @ w + b
def functional_margin(X, y, w, b):
"""
gamma_i^(f) = y_i (w^T x_i + b)
"""
return y * decision_score(X, w, b)
def geometric_margin(X, y, w, b):
"""
gamma_i = y_i (w^T x_i + b) / ||w||
"""
return functional_margin(X, y, w, b) / np.linalg.norm(w)
def project_to_hyperplane(x, w, b):
"""
Project a point x onto the hyperplane w^T x + b = 0.
"""
return x - ((w @ x + b) / np.dot(w, w)) * w
gamma_original = functional_margin(X, y, w, b)
gamma_scaled = functional_margin(X, y, w_scaled, b_scaled)
geo_original = geometric_margin(X, y, w, b)
geo_scaled = geometric_margin(X, y, w_scaled, b_scaled)
print("Original functional margins:")
print(gamma_original)
print("\nScaled functional margins:")
print(gamma_scaled)
print("\nOriginal geometric margins:")
print(geo_original)
print("\nScaled geometric margins:")
print(geo_scaled)
fig, axes = plt.subplots(1, 3, figsize=(18, 5))
x1_grid = np.linspace(-0.2, 1.8, 200)
# Since x1 + x2 - 1 = 0,
# x2 = 1 - x1
x2_boundary = 1 - x1_grid
def plot_margin_panel(ax, w_plot, b_plot, title):
scores = decision_score(X, w_plot, b_plot)
gamma = functional_margin(X, y, w_plot, b_plot)
# Decision boundary
ax.plot(
x1_grid,
x2_boundary,
linewidth=2.5,
label=r"Decision boundary: $w^\top x+b=0$"
)
# Positive and negative regions
ax.text(1.35, 1.35, r"$w^\top x+b>0$", fontsize=12)
ax.text(0.05, 0.05, r"$w^\top x+b<0$", fontsize=12)
# Plot points
for i, x in enumerate(X):
if gamma[i] > 0:
edgecolor = "green"
status = "correct"
elif gamma[i] < 0:
edgecolor = "red"
status = "wrong"
else:
edgecolor = "black"
status = "boundary"
marker = "o" if y[i] == 1 else "s"
ax.scatter(
x[0], x[1],
s=120,
marker=marker,
facecolors="white",
edgecolors=edgecolor,
linewidths=2,
zorder=5
)
ax.text(
x[0] + 0.03,
x[1] + 0.03,
f"{point_names[i]}\n$y={y[i]}$\n$\\gamma^f={gamma[i]:.1f}$",
fontsize=9
)
# Draw perpendicular projection to the boundary
x_proj = project_to_hyperplane(x, w, b)
ax.plot(
[x[0], x_proj[0]],
[x[1], x_proj[1]],
linestyle=":",
linewidth=1.5,
color="gray"
)
# Draw normal vector
x0 = np.array([0.5, 0.5]) # point on boundary
normal_unit = w_plot / np.linalg.norm(w_plot)
ax.arrow(
x0[0],
x0[1],
0.25 * normal_unit[0],
0.25 * normal_unit[1],
head_width=0.04,
length_includes_head=True,
linewidth=2,
color="black"
)
ax.text(
x0[0] + 0.2,
x0[1] + 0.22,
r"$w$",
fontsize=14
)
ax.set_title(title)
ax.set_xlabel(r"$x_1$")
ax.set_ylabel(r"$x_2$")
ax.set_xlim(-0.2, 1.8)
ax.set_ylim(-0.2, 1.8)
ax.grid(alpha=0.3)
ax.set_aspect("equal")
ax.legend(loc="upper right")
# Left panel: original margin
plot_margin_panel(
axes[0],
w,
b,
r"Original: $\gamma_i^{(f)}=y_i(w^\top x_i+b)$"
)
# Middle panel: scaled margin
plot_margin_panel(
axes[1],
w_scaled,
b_scaled,
r"Scaled: $w'=3w,\ b'=3b$"
)
# Right panel: bar chart comparing margins
bar_width = 0.25
positions = np.arange(len(point_names))
axes[2].bar(
positions - bar_width,
gamma_original,
width=bar_width,
label=r"Original functional margin"
)
axes[2].bar(
positions,
gamma_scaled,
width=bar_width,
label=r"Scaled functional margin"
)
axes[2].bar(
positions + bar_width,
geo_original,
width=bar_width,
label=r"Geometric margin"
)
axes[2].axhline(0, linewidth=1)
axes[2].set_xticks(positions)
axes[2].set_xticklabels(point_names)
axes[2].set_ylabel("Margin value")
axes[2].set_title("Functional margin changes, geometric margin does not")
axes[2].grid(axis="y", alpha=0.3)
axes[2].legend()
fig.suptitle(
"Functional Margin: Correctness, Misclassification, Boundary, and Scale Dependence",
fontsize=15,
y=1.05
)
plt.tight_layout()
plt.show()