From b5deaaffb8f2684bcbe3e500febb720886ca130b Mon Sep 17 00:00:00 2001 From: Yu-Ting Hsiung Date: Sun, 11 Oct 2026 13:03:41 +0800 Subject: [PATCH] Fix MLlib Vector dot typing --- dev/lint-python | 1 + python/pyspark/mllib/classification.py | 10 ++-- python/pyspark/mllib/linalg/__init__.py | 31 ++++++++++++- python/pyspark/mllib/regression.py | 2 +- .../mllib/tests/typing/test_linalg.yml | 46 +++++++++++++++++++ 5 files changed, 81 insertions(+), 9 deletions(-) create mode 100644 python/pyspark/mllib/tests/typing/test_linalg.yml diff --git a/dev/lint-python b/dev/lint-python index fb785c13301df..61a1b4496bc2c 100755 --- a/dev/lint-python +++ b/dev/lint-python @@ -160,6 +160,7 @@ function mypy_data_test { python/pyspark/tests/typing \ python/pyspark/sql/tests/typing \ python/pyspark/ml/tests/typing \ + python/pyspark/mllib/tests/typing \ ) 2>&1) PYTEST_STATUS=$? diff --git a/python/pyspark/mllib/classification.py b/python/pyspark/mllib/classification.py index 0d0efa9fba9b0..b0d3971ec45db 100644 --- a/python/pyspark/mllib/classification.py +++ b/python/pyspark/mllib/classification.py @@ -246,7 +246,7 @@ def predict( x = _convert_to_vector(x) if self.numClasses == 2: - margin = self.weights.dot(x) + self._intercept # type: ignore[attr-defined] + margin = self.weights.dot(x) + self._intercept if margin > 0: prob = 1 / (1 + exp(-margin)) else: @@ -272,7 +272,7 @@ def predict( best_class = i + 1 else: for i in range(0, self._numClasses - 1): - margin = x.dot(self._weightsMatrix[i]) # type: ignore[attr-defined] + margin = x.dot(self._weightsMatrix[i]) if margin > max_margin: max_margin = margin best_class = i + 1 @@ -599,7 +599,7 @@ def predict( return x.map(lambda v: self.predict(v)) x = _convert_to_vector(x) - margin = self.weights.dot(x) + self.intercept # type: ignore[attr-defined] + margin = self.weights.dot(x) + self.intercept if self._threshold is None: return margin else: @@ -799,9 +799,7 @@ def predict( if isinstance(x, RDD): return x.map(lambda v: self.predict(v)) x = _convert_to_vector(x) - return self.labels[ - numpy.argmax(self.pi + x.dot(self.theta.transpose())) # type: ignore[attr-defined] - ] + return self.labels[numpy.argmax(self.pi + x.dot(self.theta.transpose()))] def save(self, sc: SparkContext, path: str) -> None: """ diff --git a/python/pyspark/mllib/linalg/__init__.py b/python/pyspark/mllib/linalg/__init__.py index 9f19c8fbe2f1c..001f2374de2a6 100644 --- a/python/pyspark/mllib/linalg/__init__.py +++ b/python/pyspark/mllib/linalg/__init__.py @@ -327,6 +327,21 @@ def asML(self) -> newlinalg.Vector: """ raise NotImplementedError + @overload + def dot(self, other: Union["Vector", List[float], Tuple[float, ...], range]) -> np.float64: ... + + @overload + def dot(self, other: "VectorLike") -> Union[np.float64, np.ndarray]: ... + + def dot(self, other: "VectorLike") -> Union[np.float64, np.ndarray]: + """ + Compute the dot product with a vector or matrix. + + Vector operands return a scalar. NumPy and SciPy operands may + return an array, depending on their dimensions and representation. + """ + raise NotImplementedError + def __len__(self) -> int: raise NotImplementedError @@ -414,7 +429,13 @@ def norm(self, p: "NormType") -> np.floating[Any]: """ return np.linalg.norm(self.array, p) - def dot(self, other: "VectorLike") -> np.float64: + @overload + def dot(self, other: Union[Vector, List[float], Tuple[float, ...], range]) -> np.float64: ... + + @overload + def dot(self, other: "VectorLike") -> Union[np.float64, np.ndarray]: ... + + def dot(self, other: "VectorLike") -> Union[np.float64, np.ndarray]: """ Compute the dot product of two Vectors. We support (Numpy array, list, SparseVector, or SciPy sparse) @@ -765,7 +786,13 @@ def parse(s: str) -> "SparseVector": raise ValueError("Unable to parse values from %s." % s) return SparseVector(cast(int, size), indices, values) - def dot(self, other: "VectorLike") -> np.float64: + @overload + def dot(self, other: Union[Vector, List[float], Tuple[float, ...], range]) -> np.float64: ... + + @overload + def dot(self, other: "VectorLike") -> Union[np.float64, np.ndarray]: ... + + def dot(self, other: "VectorLike") -> Union[np.float64, np.ndarray]: """ Dot product with a SparseVector or 1- or 2-dimensional Numpy array. diff --git a/python/pyspark/mllib/regression.py b/python/pyspark/mllib/regression.py index 7fb5ff055da68..540f069a3bf3b 100644 --- a/python/pyspark/mllib/regression.py +++ b/python/pyspark/mllib/regression.py @@ -162,7 +162,7 @@ def predict(self, x: Union["VectorLike", RDD["VectorLike"]]) -> Union[float, RDD if isinstance(x, RDD): return x.map(self.predict) x = _convert_to_vector(x) - return self.weights.dot(x) + self.intercept # type: ignore[attr-defined] + return self.weights.dot(x) + self.intercept @inherit_doc diff --git a/python/pyspark/mllib/tests/typing/test_linalg.yml b/python/pyspark/mllib/tests/typing/test_linalg.yml new file mode 100644 index 0000000000000..9fb0f3d7e7201 --- /dev/null +++ b/python/pyspark/mllib/tests/typing/test_linalg.yml @@ -0,0 +1,46 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +- case: mllibVectorDot + main: | + from typing import Union + from typing_extensions import assert_type + import numpy as np + from pyspark.mllib._typing import VectorLike + from pyspark.mllib.linalg import DenseVector, SparseVector, Vector + + def check_dot( + vector: Vector, + dense: DenseVector, + sparse: SparseVector, + array: np.ndarray, + other: VectorLike, + ) -> None: + assert_type(vector.dot(vector), np.float64) + assert_type(dense.dot(vector), np.float64) + assert_type(sparse.dot(vector), np.float64) + assert_type(vector.dot([1.0, 2.0]), np.float64) + assert_type(vector.dot((1.0, 2.0)), np.float64) + assert_type(vector.dot(range(2)), np.float64) + # An ndarray's type does not distinguish vector and matrix operands. + assert_type(vector.dot(array), Union[np.float64, np.ndarray]) + assert_type(dense.dot(array), Union[np.float64, np.ndarray]) + assert_type(sparse.dot(array), Union[np.float64, np.ndarray]) + # VectorLike also includes SciPy sparse matrices and arrays. + assert_type(vector.dot(other), Union[np.float64, np.ndarray]) + assert_type(dense.dot(other), Union[np.float64, np.ndarray]) + assert_type(sparse.dot(other), Union[np.float64, np.ndarray])