From de81f95bad2b7ea02448b7650d3ecac9dfbff160 Mon Sep 17 00:00:00 2001 From: James Nelson Date: Tue, 28 Jul 2026 15:04:32 +0100 Subject: [PATCH] regression --- qkernel/model.py | 60 +++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 59 insertions(+), 1 deletion(-) diff --git a/qkernel/model.py b/qkernel/model.py index 06ff038..45ac625 100644 --- a/qkernel/model.py +++ b/qkernel/model.py @@ -2,7 +2,7 @@ import numpy as np from sklearn.svm import SVC - +from sklearn.kernel_ridge import KernelRidge class Model(ABC): """ @@ -113,3 +113,61 @@ def predict(self, gram: np.ndarray) -> np.ndarray: Predicted labels. """ return self.model.predict(gram) + + +class KRR(Model): + """ + Quantum-kernel-compatible Kernel Ridge Regression wrapper. + + This class uses scikit-learn's KernelRidge with a precomputed kernel, + making it suitable for quantum kernel methods (QKRR-style pipelines). + + Parameters + ---------- + alpha : float, default=1.0 + Regularization strength. Larger values mean stronger regularization. + **kwargs + Additional keyword arguments passed to sklearn.kernel_ridge.KernelRidge. + """ + + def __init__(self, alpha: float = 1.0, **kwargs) -> None: + """ + Initialize the KRR model with a precomputed kernel. + + Parameters + ---------- + alpha : float, default=1.0 + Regularization strength. + **kwargs + Keyword arguments forwarded to sklearn's KernelRidge. + """ + self.model = KernelRidge(kernel="precomputed", alpha=alpha, **kwargs) + + def train(self, gram_train: np.ndarray, y_train: np.ndarray) -> None: + """ + Train the KRR model. + + Parameters + ---------- + gram_train : np.ndarray + Training Gram matrix of shape (n_train, n_train). + y_train : np.ndarray + Training targets. + """ + self.model.fit(gram_train, y_train) + + def predict(self, gram: np.ndarray) -> np.ndarray: + """ + Predict target values. + + Parameters + ---------- + gram : np.ndarray + Gram matrix of shape (n_samples, n_train) for prediction. + + Returns + ------- + np.ndarray + Predicted values. + """ + return self.model.predict(gram)