Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion machine_learning/k_means_clust.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,10 +48,12 @@
"""

import warnings
from typing import cast

import numpy as np
import pandas as pd
from matplotlib import pyplot as plt
from mpl_toolkits.mplot3d import Axes3D
from sklearn.metrics import pairwise_distances

warnings.filterwarnings("ignore")
Expand Down Expand Up @@ -157,7 +159,10 @@ def plot_heterogeneity(heterogeneity, k) -> None:


def plot_kmeans(data, centroids, cluster_assignment) -> None:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We need type hints for data, centroids, and cluster_assignment

ax = plt.axes(projection="3d")
# plt.axes() is typed to return the base 2D Axes, but projection="3d" makes
# it an Axes3D at runtime; cast so ty resolves Axes3D.scatter's (xs, ys, zs)
# signature instead of colliding its 3rd positional arg with the `s` kwarg.
ax = cast(Axes3D, plt.axes(projection="3d"))
ax.scatter(data[:, 0], data[:, 1], data[:, 2], c=cluster_assignment, cmap="viridis")
ax.scatter(
centroids[:, 0], centroids[:, 1], centroids[:, 2], c="red", s=100, marker="x"
Expand Down
1 change: 0 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -323,7 +323,6 @@ rules.invalid-return-type = "ignore"
rules.no-matching-overload = "ignore"
rules.not-iterable = "ignore"
rules.not-subscriptable = "ignore"
rules.parameter-already-assigned = "ignore"
rules.unresolved-attribute = "ignore"
rules.unresolved-import = "ignore"
rules.unsupported-operator = "ignore"
Expand Down
Loading