from contextlib import contextmanager
import copy
import numpy as np
from skbase.base import BaseMetaObject
from pyqit.core.trainer import Trainer, console, has_rich
from pyqit.data.datamodule import DataModule, _apply_prescale
from pyqit.utils.utils import (
_cat,
_count_params,
_ensure_col,
_is_classifier,
_is_torch,
_mean,
_stack,
_to_numpy,
)
[docs]
class PipelineStage:
"""One model in a `QuantumPipeline`.
Parameters
----------
model : BaseModel
name : str, optional
Defaults to the model's class name.
passthrough : bool, default False
Concatenate this stage's input onto its output.
trainable : bool, default True
`frozen_backbone` fit mode requires this False on every non-final stage.
input_slice : slice or array-like of int, optional
Feed only these input columns to this stage.
"""
def __init__(
self, model, name=None, passthrough=False, trainable=True, input_slice=None
):
self.model = model
self.name = name or type(model).__name__
self.passthrough = passthrough
self.trainable = trainable
self.input_slice = input_slice
def _flags(self):
return [
flag
for flag, on in (
("frozen", not self.trainable),
("passthrough", self.passthrough),
(f"slice={self.input_slice}", self.input_slice is not None),
)
if on
]
def __repr__(self):
flags = self._flags()
flag_str = f" [{', '.join(flags)}]" if flags else ""
return f"Stage({self.name}{flag_str})"
[docs]
class QuantumPipeline(BaseMetaObject):
"""Compose `PipelineStage` objects sequentially or as an ensemble.
Parameters
----------
steps : list of PipelineStage, or list of (name, model)
mode : str, default "sequential"
How the stages relate to each other, which decides what each stage
receives as input.
- ``"sequential"``: each stage's output is the next stage's input, and
the last stage's output is the pipeline's.
- ``"ensemble"``: every stage receives the same input and the outputs
are combined by `aggregation`.
aggregation : str or callable, default "mean"
How ensemble outputs are combined. Ignored in sequential mode.
- ``"mean"``: averages the stage outputs.
- ``"vote"``: takes the majority of the stages' hard labels.
- callable: receives the list of raw stage outputs and returns the
combined output.
fit_mode : str, default "sequential_greedy"
How the stages of a sequential pipeline are trained, which decides
what loss each stage sees. Ignored in ensemble mode, where every
trainable stage trains independently on the same data.
- ``"sequential_greedy"``: trains each stage in turn against the
labels, on the output of the stages before it.
- ``"frozen_backbone"``: trains only the final stage. Every other
stage must have ``trainable=False``.
- ``"joint"``: trains every trainable stage at once against the final
stage's loss. Gradients pass through every stage, frozen ones
included. Use this for hybrid classical-quantum networks.
Examples
--------
>>> from pyqit.core import PipelineStage, QuantumPipeline
>>> pipe = QuantumPipeline(
... [PipelineStage(backbone, trainable=False), PipelineStage(head)],
... fit_mode="frozen_backbone",
... )
>>> pyqit.Trainer(max_epochs=20).fit(pipe, dm) # doctest: +SKIP
A hybrid network, dense to circuit to dense, trained end to end:
>>> from pyqit.models.layers import DenseClassifier, DenseLayer, QuantumLayer
>>> hybrid = QuantumPipeline(
... [
... DenseLayer(n_features=8, n_out=4, activation="tanh"),
... QuantumLayer(n_qubits=4, n_layers=2),
... DenseClassifier(n_features=4),
... ],
... fit_mode="joint",
... )
>>> history = pyqit.Trainer(max_epochs=20).fit(hybrid, dm) # doctest: +SKIP
"""
_tags = {
"object_type": "pipeline",
"authors": "phoeenniixx",
"mode": "sequential",
"n_stages": 0,
"has_quantum": False,
}
def __init__(
self, steps, mode="sequential", aggregation="mean", fit_mode="sequential_greedy"
):
if mode not in ("sequential", "ensemble"):
raise ValueError(f"mode must be 'sequential' or 'ensemble', got {mode}.")
if fit_mode not in ("sequential_greedy", "frozen_backbone", "joint"):
raise ValueError(
"fit_mode must be 'sequential_greedy', 'frozen_backbone' or "
f"'joint', got {fit_mode}."
)
self.steps = self._to_named_steps(steps)
self.mode = mode
self.aggregation = aggregation
self.fit_mode = fit_mode
super().__init__()
has_quantum = any(
s.model.get_tag("is_quantum", False, raise_error=False)
for _, s in self.steps
)
kinds = [
s.model.get_tag("estimator_type", None, raise_error=False)
for _, s in self.steps
]
if mode == "ensemble" and len(set(kinds)) > 1:
raise ValueError(
"Ensemble stages must all be classifiers or all regressors, since "
f"their outputs are aggregated into one prediction; got {kinds}."
)
self.set_tags(
mode=mode,
n_stages=len(self.steps),
has_quantum=has_quantum,
estimator_type=kinds[-1],
)
[docs]
def set_params(self, **kwargs):
"""Set stage or nested-stage parameters. See `sklearn`'s convention."""
return self._set_params("steps", **kwargs)
[docs]
def get_params(self, deep: bool = True):
"""Get stage and nested-stage parameters. See `sklearn`'s convention."""
return self._get_params("steps", deep=deep)
@property
def named_stages(self) -> dict[str, PipelineStage]:
"""Stages keyed by name."""
return dict(self.steps)
def __getitem__(self, key: str | int) -> PipelineStage:
if isinstance(key, int):
return self.steps[key][1]
return self.named_stages[key]
def __len__(self) -> int:
return len(self.steps)
@staticmethod
def _slice_input(X, input_slice):
"""Select ``input_slice`` columns from ``X``, always returning 2-D."""
if X is None or input_slice is None:
return X
if isinstance(input_slice, (int, np.integer)):
return X[:, [int(input_slice)]]
return X[:, input_slice]
@staticmethod
def _prescale_of(model):
return getattr(type(getattr(model, "embedding_obj", None)), "PRESCALE", None)
def _prescale_for(self, stage, X):
"""Shape ``X`` for ``stage``'s embedding, as ``DataModule.setup`` does."""
n_qubits = getattr(stage.model, "n_qubits", None)
prescale = self._prescale_of(stage.model)
if X is None or n_qubits is None:
return X
return _apply_prescale(X, prescale, n_qubits)
def _prepare_stage_input(self, X, stage):
"""Slice, then prescale, the input a stage is about to consume.
Every stage is prescaled here, the first included: the pipeline sets its
DataModule up without an encoder, so slicing and passthrough see the
unprescaled features and each embedding's prescaling runs exactly once.
"""
return self._prescale_for(stage, self._slice_input(X, stage.input_slice))
def _run_stage(self, stage, X, **weights):
"""Run one non-final stage on ``X``; return what the next one receives.
Passthrough concatenates the unprescaled slice, not the prescaled input,
so the next stage's prescaling is applied to it once rather than twice.
"""
raw = self._slice_input(X, stage.input_slice)
inp = self._prescale_for(stage, raw)
out = _ensure_col(stage.model.forward(inp, **weights))
return _cat(raw, out) if stage.passthrough else out
def _run_sequential(self, X, labels=False, **custom_weights):
current = X if _is_torch(X) else np.asarray(X)
for name, stage in self.steps[:-1]:
current = self._run_stage(
stage, current, **self._stage_weights(name, custom_weights)
)
name, head = self.steps[-1]
inp = self._prepare_stage_input(current, head)
if labels:
return head.model.predict_step(inp)
return head.model.forward(inp, **self._stage_weights(name, custom_weights))
@staticmethod
def _stage_weights(name, flat_weights):
prefix = f"{name}."
return {
k[len(prefix) :]: v for k, v in flat_weights.items() if k.startswith(prefix)
}
@property
def weights(self):
"""Flat ``{"<stage>.<qnode>.<weight>": array}`` dict of trainable stages."""
return {
f"{name}.{key}": value
for name, stage in self.steps
if stage.trainable
for key, value in stage.model.weights.items()
}
[docs]
def weight_groups(self) -> dict:
"""The trainable stages' `weight_groups`, keyed like `weights`."""
groups = {}
for name, stage in self.steps:
if stage.trainable:
for group, keys in stage.model.weight_groups().items():
groups.setdefault(group, []).extend(f"{name}.{k}" for k in keys)
return groups
[docs]
def update_weights(self, flat_weights_dict):
"""Write `flat_weights_dict`, keyed like `weights`, into the stages."""
for name, stage in self.steps:
own = self._stage_weights(name, flat_weights_dict)
if own:
stage.model.update_weights(own)
@property
def _qnodes(self):
return {
f"{name}__{key}".replace(".", "_"): node
for name, stage in self.steps
if stage.trainable
for key, node in getattr(stage.model, "_qnodes", {}).items()
}
def _expected_width(self, model):
"""Feature width ``model`` expects."""
n_qubits = getattr(model, "n_qubits", None)
if n_qubits is None:
return None
return 2**n_qubits if self._prescale_of(model) == "amplitude" else n_qubits
def _check_ensemble_consistent(self):
"""Ensemble stages all receive the same X, so they must agree on shape."""
sliced = [n for n, s in self.steps if s.input_slice is not None]
if sliced:
raise ValueError(
f"Stages {sliced} set input_slice, but ensemble mode feeds every "
"stage the same X and aggregates the results, so per-stage "
"column views are not applied. Use mode='sequential'."
)
specs = []
for name, stage in self.steps:
embedding = getattr(stage.model, "embedding_obj", None)
encoder = type(embedding).__name__ if embedding is not None else None
specs.append((name, self._expected_width(stage.model), encoder))
if len({s[1:] for s in specs}) > 1:
detail = "; ".join(f"{n}: width={w}, encoder={e}" for n, w, e in specs)
raise ValueError(
"Ensemble stages are all fed the same input, but they disagree "
f"on what they accept: {detail}. Give every stage the same "
"n_qubits and encoder, or use mode='sequential'."
)
[docs]
def forward(self, X, **custom_weights):
"""Run every stage on `X` and return the pipeline's raw output.
`X` is split and normalized but not prescaled: each stage's embedding
prescaling is applied here, as `predict` and `fit` do.
`custom_weights` override the stages' own weights, keyed as in
`weights`; sequential mode only.
"""
if self.mode == "sequential":
return self._run_sequential(X, **custom_weights)
return self._forward_ensemble(X)
def _forward_ensemble(self, X):
inputs = [self._prepare_stage_input(X, stage) for _, stage in self.steps]
if self.aggregation == "vote":
labels = np.stack(
[
_to_numpy(stage.model.predict_step(x)).ravel().astype(int)
for (_, stage), x in zip(self.steps, inputs)
]
)
return np.apply_along_axis(lambda v: np.bincount(v).argmax(), 0, labels)
raw = [stage.model.forward(x) for (_, stage), x in zip(self.steps, inputs)]
if callable(self.aggregation):
return self.aggregation(raw)
if self.aggregation == "mean":
return _mean(_stack(raw), axis=0)
raise ValueError(f"Unknown aggregation: {self.aggregation}")
def _fit(self, datamodule: DataModule, trainer: Trainer) -> dict:
"""Fit every trainable stage with ``trainer``.
Parameters
----------
datamodule : DataModule
Split and normalized here, never prescaled: the pipeline prescales
every stage's input itself, so a DataModule already set up for a
single model's encoder is rejected.
trainer : Trainer
Used for every trainable stage in turn.
Returns
-------
dict or TrainingHistory
`TrainingHistory` per trained stage, keyed by stage name; a single
`TrainingHistory` under `fit_mode="joint"`.
"""
if datamodule.encoder is not None:
raise ValueError(
f"This DataModule is already prescaled for "
f"{datamodule.encoder.__name__}, but a QuantumPipeline prescales "
"each stage's input itself, so it would be scaled twice. Pass a "
"DataModule that has not been set up for a single model."
)
if self.mode == "ensemble":
self._check_ensemble_consistent()
elif self.fit_mode == "frozen_backbone":
trainable = [name for name, s in self.steps[:-1] if s.trainable]
if trainable:
raise ValueError(
"frozen_backbone mode requires all stages except the last to "
f"be frozen. Stage '{trainable[0]}' has trainable=True."
)
datamodule.setup(batch_size=trainer.batch_size)
verbose = trainer.verbose
sequential = self.mode == "sequential"
self._print_pipeline_summary(verbose)
if sequential and self.fit_mode == "joint":
with self._quiet_stage_summary(trainer):
return trainer._fit_model(self, datamodule)
dm = datamodule
histories = {}
for i, (name, stage) in enumerate(self.steps):
if stage.trainable:
histories[name] = self._fit_stage(i, dm, trainer, verbose)
if sequential and i < len(self.steps) - 1:
dm = self._transform_datamodule(dm, stage)
return histories
def _fit_stage(self, idx, dm, trainer, verbose):
name, stage = self.steps[idx]
stage_dm = self._stage_datamodule(dm, stage)
self._announce_stage(verbose, idx, name)
with self._quiet_stage_summary(trainer):
return trainer.fit(stage.model, stage_dm)
def _stage_datamodule(self, dm: DataModule, stage: PipelineStage) -> DataModule:
"""Slice and prescale ``dm`` into the view ``stage`` actually consumes."""
return dm._map_features(lambda X: self._prepare_stage_input(X, stage))
def _transform_datamodule(self, dm: DataModule, stage: PipelineStage) -> DataModule:
"""Run ``stage`` over every split of the unprescaled ``dm``."""
return dm._map_features(lambda X: self._transform_split(stage, X, dm._backend))
def _transform_split(self, stage, X, backend):
if backend != "torch" and not _is_torch(X):
return self._run_stage(stage, X)
import torch
if not _is_torch(X):
X = torch.as_tensor(np.asarray(X), dtype=torch.float32)
with torch.no_grad():
out = self._run_stage(stage, X)
return out.detach() if _is_torch(out) else out
@staticmethod
def _to_named_steps(steps: list) -> list[tuple[str, PipelineStage]]:
from pyqit.models.base.base import BaseModel
named, seen = [], set()
for i, step in enumerate(steps):
name, obj = (
step if isinstance(step, tuple) and len(step) == 2 else (None, step)
)
if isinstance(obj, BaseModel):
obj = PipelineStage(obj, name=name)
if not isinstance(obj, PipelineStage) or not isinstance(
obj.model, BaseModel
):
raise TypeError(
f"Each stage must be a BaseModel, PipelineStage, or "
f"(name, model/stage) tuple. Got {type(step).__name__} at "
f"index {i}."
)
base = name = name or obj.name
counter = 1
while name in seen:
name = f"{base}_{counter}"
counter += 1
seen.add(name)
obj.name = name
named.append((name, obj))
return named
[docs]
def predict_step(self, X):
"""Run every stage on `X`, hard-labeling the final stage's output.
A regressor's output is returned as it is.
"""
if self.mode == "sequential":
return self._run_sequential(X, labels=True)
out = self._forward_ensemble(X)
if self.aggregation == "vote" or not _is_classifier(self):
return out
if out.ndim > 1 and out.shape[1] > 1:
return out.argmax(1)
labels = out >= 0.5
return labels.int() if _is_torch(labels) else labels.astype(int)
@staticmethod
@contextmanager
def _quiet_stage_summary(trainer):
"""Suppress a stage trainer's per-model summary for the duration of a fit.
The pipeline prints one summary that names every stage; the trainer's own
table would repeat it once per stage without saying which stage it is.
Restored afterwards so a trainer reused outside the pipeline is unchanged.
"""
previous = getattr(trainer, "_print_summary", True)
trainer._print_summary = False
try:
yield
finally:
trainer._print_summary = previous
def _print_pipeline_summary(self, verbose: int):
"""Print the pipeline structure once, in place of N per-model tables."""
if verbose < 2:
return
header = f"QuantumPipeline | mode={self.mode} | stages={len(self.steps)}"
if self.mode == "ensemble":
header += f" | aggregation={self.aggregation}"
else:
header += f" | fit_mode={self.fit_mode}"
rows = [
(
str(i + 1),
name,
type(stage.model).__name__,
str(getattr(stage.model, "n_qubits", "N/A")),
str(_count_params(stage.model) or 0),
", ".join(stage._flags()) or "trainable",
)
for i, (name, stage) in enumerate(self.steps)
]
columns = ("#", "Stage", "Model", "Qubits", "Params", "Status")
if not has_rich():
print(f"\n[Pipeline] {header}")
widths = [
max(len(c), *(len(r[j]) for r in rows)) for j, c in enumerate(columns)
]
fmt = " ".join(f"{{:<{w}}}" for w in widths)
print(" " + fmt.format(*columns))
for row in rows:
print(" " + fmt.format(*row))
print()
return
from rich.table import Table
table = Table(show_header=True, header_style="bold cyan", box=None, title=None)
for column in columns:
table.add_column(column, style="bold" if column == "Stage" else None)
for row in rows:
table.add_row(*row)
console().print(f"[bold cyan][Pipeline][/bold cyan] {header}")
console().print(table)
console().print()
def _announce_stage(self, verbose: int, idx: int, name: str):
"""Name the stage about to train, since its trainer no longer does."""
if verbose < 1:
return
label = f"[{idx + 1}/{len(self.steps)}] Fitting stage {name!r}"
if has_rich():
console().print(f"[bold cyan]{label}[/bold cyan]")
else:
print(label)
[docs]
def clone(self) -> "QuantumPipeline":
"""Return an independent copy: stages deep-copied, weights included."""
return type(self)(
steps=[(name, copy.deepcopy(stage)) for name, stage in self.steps],
mode=self.mode,
aggregation=self.aggregation,
fit_mode=self.fit_mode,
)
def __call__(self, X):
"""Alias for `forward`."""
return self.forward(X)
def __repr__(self):
stage_str = "\n ".join(str(s) for s in self.steps)
return f"QuantumPipeline(mode={self.mode})\n {stage_str}"