Skip to content
Draft
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
60 changes: 59 additions & 1 deletion qkernel/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@

import numpy as np
from sklearn.svm import SVC

from sklearn.kernel_ridge import KernelRidge

class Model(ABC):
"""
Expand Down Expand Up @@ -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)