"""Qubit-count scaling projections for resource usage."""
from __future__ import annotations
from itertools import combinations
from typing import Any, Sequence
from .._logging import get_logger
from .estimator import ResourceReport
log = get_logger(__name__)
def _n_filters(num_qubits: int, num_levels: int, n_filter_types: int = 1) -> int:
"""Number of filter features for a given qubit/level configuration."""
pairs_per_qubit = len(list(combinations(range(num_levels), 2)))
return pairs_per_qubit * num_qubits * n_filter_types
[docs]
def project_scalability(
base_report: ResourceReport,
qubit_counts: Sequence[int],
*,
base_num_qubits: int = 5,
num_levels: int = 2,
n_filter_types: int = 1,
) -> list[dict[str, Any]]:
"""Project resource usage to different qubit counts.
Extrapolates from a base measurement or estimate using known scaling
relationships:
- Filter count scales as ``C(num_levels, 2) * num_qubits * n_filter_types``.
- NN input dim scales linearly with filter count.
- NN compute (MACs) scales quadratically with input dim for the first
hidden layer, linearly for subsequent layers.
- Threshold/comparator resources scale linearly.
Args:
base_report: Measured or estimated resources at *base_num_qubits*.
qubit_counts: Target qubit counts to project to.
base_num_qubits: Qubit count the *base_report* was measured at.
num_levels: Energy levels per qubit.
n_filter_types: Number of filter families (1=MF, 2=MF+RMF, 3=MF+RMF+EMF).
Returns:
List of dicts with ``"num_qubits"``, ``"lut"``, ``"ff"``,
``"bram"``, ``"dsp"``, ``"latency_cycles"``, ``"latency_ns"``,
and ``"scale_factor"``.
"""
base_filters = _n_filters(base_num_qubits, num_levels, n_filter_types)
if base_filters == 0:
base_filters = 1
results: list[dict[str, Any]] = []
for nq in qubit_counts:
target_filters = _n_filters(nq, num_levels, n_filter_types)
linear_scale = target_filters / base_filters
quadratic_scale = linear_scale**2
_nn_types = {
"fnn", "hybrid", "cnn", "transformer", "herqules",
"leakage", "multilevel", "mlp",
}
is_nn = base_report.classifier_name in _nn_types
scale = quadratic_scale if is_nn else linear_scale
lat_scale = max(1.0, linear_scale**0.5) if is_nn else linear_scale
def _scale_val(val: int | float | None, s: float) -> int | None:
return int(val * s) if val is not None else None
lat_cycles = base_report.latency_cycles
if lat_cycles is not None:
lat_cycles = int(lat_cycles * lat_scale)
lat_ns = base_report.latency_ns
if lat_ns is not None:
lat_ns = float(lat_ns * lat_scale)
proj = {
"num_qubits": nq,
"lut": _scale_val(base_report.lut, scale),
"ff": _scale_val(base_report.ff, scale),
"bram": _scale_val(base_report.bram, max(1.0, linear_scale)),
"dsp": _scale_val(base_report.dsp, scale),
"latency_cycles": lat_cycles,
"latency_ns": lat_ns,
"scale_factor": round(scale, 3),
"filter_count": target_filters,
}
results.append(proj)
log.info(
"Projected %s to %d qubit counts (base=%d qubits, %d filters)",
base_report.classifier_name,
len(qubit_counts),
base_num_qubits,
base_filters,
)
return results
# Public alias
project_qubit_scaling = project_scalability