TORX

Simulators

A simulator compiles a circuit and reads it back as samples, moments, or a density.

Abstract base classes

AbstractCompiledPCircuitclass
AbstractCompiledPCircuit()

Abstract parent class for probabilistic circuits built for a specific backend.

from_pcircuitclassmethod
from_pcircuit(circuit: ~_CircuitType, thetas: PyTree[Array]) -> Self

Compile the given probabilistic circuit with the given parameters.

Arguments:

  • circuit: The probabilistic circuit to compile
  • thetas: The parameters to bake into the compiled circuit.

Returns:

The compiled circuit.

to_pcircuitmethod
to_pcircuit(structure: ~_CircuitType) -> ~_CircuitType

Return a circuit structure reconstructed by this compiled backend.

AbstractSimulatorclass
AbstractSimulator()

Abstract parent class for probabilistic circuit simulators.

expvalmethod
expval(*args, **kwargs)

Compute the expectation value of the given index after circuit execution.

expval_allmethod
expval_all(*args, **kwargs)

Compute the expectation value of all indices after circuit execution.

build_circuitmethod
build_circuit(circuit: ~_CircuitType, thetas: PyTree[Array]) -> ~_CompiledType

Build the circuit for this simulator with the given parameters.

Arguments:

  • circuit: The circuit to build
  • thetas: The parameters to bake into the compiled circuit.

Returns:

The built circuit.

Concrete classes

GaussianMomentsclass
GaussianMoments(
    mean: Float[Array, 'continuous_dim'],
    covariance: Float[Array, 'continuous_dim continuous_dim'],
    site_offsets: tuple[tuple[int, int, int], ...],
    observed_sites: tuple[int, ...] = (),
)

Joint Gaussian moments over continuous sites.

mean and covariance follow the order given by sites. site_offsets maps each site to its (start, stop) slice into them.

meanattribute
mean: <class 'jaxFloat[Array, 'continuous_dim']'>
covarianceattribute
covariance: <class 'jaxFloat[Array, 'continuous_dim continuous_dim']'>
site_offsetsattribute
site_offsets: tuple[tuple[int, int, int], ...]
observed_sitesattribute
observed_sites: tuple[int, ...]
sitesproperty
sites

Sites covered by this state, in mean/covariance order.

site_indicesmethod
site_indices(site: int) -> Int[Array, 'd']

Return flat state-vector indices for one continuous site.

site_momentsmethod
site_moments(site: int) -> tuple[Float[Array, 'd'], Float[Array, 'd d']]

Return marginal mean and covariance for one continuous site.

CompiledAffineGaussianPCircuitclass
CompiledAffineGaussianPCircuit(
    gates: list[AbstractDiscreteGate | AbstractHybridGate],
    thetas: list[PyTree[Array]],
    site_offsets: tuple[tuple[int, int, int], ...],
    reps: int,
)

Compiled hybrid circuit for the affine Gaussian simulator.

gatesattribute
gates: list[AbstractDiscreteGate | AbstractHybridGate]
thetasattribute
thetas: list[jaxPyTree[Array]]
site_offsetsattribute
site_offsets: tuple[tuple[int, int, int], ...]
repsattribute
reps: int
from_pcircuitclassmethod
from_pcircuit(circuit: HybridPCircuit, thetas: list[PyTree[Array]]) -> Self

Compile circuit for exact affine Gaussian moment propagation.

Arguments:

  • circuit: The hybrid circuit to compile.
  • thetas: Per-gate parameters aligned with circuit.gates.

Returns:

The compiled circuit.

to_pcircuitmethod
to_pcircuit(structure: HybridPCircuit) -> HybridPCircuit

Return a HybridPCircuit with structure's gates and compiled reps.

AffineGaussianSimulatorclass
AffineGaussianSimulator()

Exact moment simulator for the affine Gaussian fragment of hybrid circuits.

Propagates the joint mean and covariance in closed form through each gate's affine Gaussian channel (A, b, log_var), exposed via [AbstractAffineGaussianGate][torx.psc.AbstractAffineGaussianGate], and conditions on observed sites via the Schur complement. Only supports Gaussian gates.

Dense reference implementation: O(D^3) per gate in the total continuous dimension D; intended for small affine-Gaussian circuits.

build_circuitmethod
build_circuit(
    circuit: HybridPCircuit,
    thetas: list[PyTree[Array]],
) -> CompiledAffineGaussianPCircuit

Compile circuit for exact affine Gaussian moment propagation.

propagatemethod
propagate(
    circuit: CompiledAffineGaussianPCircuit,
    initial_continuous: Float[Array, 'continuous_dim'],
) -> GaussianMoments

Propagate joint Gaussian moments through an affine-Gaussian circuit.

Arguments:

  • circuit: The compiled affine Gaussian circuit to execute.
  • initial_continuous: Flat initial continuous state.

Returns:

GaussianMoments over every continuous site.

conditionmethod
condition(
    circuit: CompiledAffineGaussianPCircuit,
    observations: Mapping[int, Array] | None = None,
    *,
    initial_continuous: Float[Array, 'continuous_dim'],
    query_sites: Sequence[int] | None = None,
    jitter: Float[Array, ''] | float = 0.0,
) -> GaussianMoments

Condition final Gaussian moments on continuous-site observations.

The circuit is propagated to a joint Gaussian over the final continuous state, then the queried sites are conditioned on the observed ones via the Schur complement. The solve uses a Cholesky factorization of the observed covariance plus jitter on the diagonal. Observed/query site membership must be static Python-level dict keys or Sequence entries; only observation values and jitter may be traced, so JIT/vmap over which sites are observed is unsupported by design.

Arguments:

  • circuit: The compiled affine Gaussian circuit to execute.
  • observations: Mapping from each observed continuous site to its
  • value.

  • initial_continuous: Flat initial continuous state.
  • query_sites: Sites to return. If None, all unobserved sites are
  • queried.

  • jitter: Nonnegative diagonal regularizer for the observed
  • covariance.

Returns:

GaussianMoments over the queried sites, with observed_sites set.

expvalmethod
expval(
    circuit: CompiledAffineGaussianPCircuit,
    initial_continuous: Float[Array, 'continuous_dim'],
    site: int = 0,
) -> Float[Array, 'd']

Return the marginal mean of one continuous site.

Arguments:

  • circuit: The compiled affine Gaussian circuit to execute.
  • initial_continuous: Flat initial continuous state.
  • site: The continuous site whose marginal mean to return.

Returns:

The marginal mean of site.

expval_allmethod
expval_all(
    circuit: CompiledAffineGaussianPCircuit,
    initial_continuous: Float[Array, 'continuous_dim'],
) -> Float[Array, 'continuous_dim']

Return the joint mean over all continuous sites.

CompiledBranchingPCircuitclass
CompiledBranchingPCircuit(
    num_pdits: int,
    reps: int,
    max_branches: int,
    branch_ops: Int[Array, 'num_gates max_branches max_basis max_l'],
    num_branches: Int[Array, 'num_gates'],
    sites: Int[Array, 'num_gates l'],
    dims: Int[Array, 'num_gates l'],
    basis_sizes: Int[Array, 'num_gates'],
    thetas: Float[Array, 'num_gates max_branches_minus_1'],
)

Compiled probabilistic circuit for the branching simulator.

num_pditsattribute
num_pdits: int
repsattribute
reps: int
max_branchesattribute
max_branches: int
branch_opsattribute
branch_ops: <class 'jaxInt[Array, 'num_gates max_branches max_basis max_l']'>
num_branchesattribute
num_branches: <class 'jaxInt[Array, 'num_gates']'>
sitesattribute
sites: <class 'jaxInt[Array, 'num_gates l']'>
dimsattribute
dims: <class 'jaxInt[Array, 'num_gates l']'>
basis_sizesattribute
basis_sizes: <class 'jaxInt[Array, 'num_gates']'>
thetasattribute
thetas: <class 'jaxFloat[Array, 'num_gates max_branches_minus_1']'>
from_pcircuitclassmethod
from_pcircuit(circuit: DiscretePCircuit, thetas: list[Float[Array, '...']]) -> Self

Compile the given probabilistic circuit with the given parameters.

This compiled form removes the list of probabilistic gates in favour of a more JIT-friendly representation. This representation uses the matrix forms of the branches of the probabilistic gates.

Arguments:

  • circuit: The probabilistic circuit to compile
  • thetas: Per-gate parameters aligned with circuit.gates; stacked and padded to (num_gates, max_branches - 1) with -inf.

Returns:

The compiled circuit.

to_pcircuitmethod
to_pcircuit(structure: DiscretePCircuit) -> DiscretePCircuit

Return a DiscretePCircuit with structure's gates and compiled reps.

to_thetasmethod
to_thetas() -> list[Float[Array, '...']]

Recover the per-gate parameter list from the padded thetas.

Inverse of the stacking done in from_pcircuit: each gate's theta is the first K - 1 entries of its padded row.

Returns:

A list of per-gate thetas, aligned with the original gates.

BranchingSimulatorclass
BranchingSimulator(
    diff_method: Literal['param_shift_inf', 'param_shift_single', 'param_shift_filter'] = 'param_shift_inf',
    num_samples: int = 1,
)

A branch-sampling simulator for lookup-table probabilistic circuits.

Instead of storing full state vectors, it samples from distributions.

Choosing a differentiation method

Three differentiation methods are available:

  • "param_shift_inf": Uses the parameter shift rule with deterministic
  • gates ($\theta \to \pm\infty$). Requires $2N$ circuit evaluations for $N$ parameters.

  • "param_shift_single": Uses the parameter shift rule with primal reuse.
  • Requires $N$ circuit evaluations.

  • "param_shift_filter": Estimates gradients from a single forward pass
  • by filtering samples based on which branch was taken at each gate. For each gate, samples are partitioned into those that applied the gate vs. those that did not, and the gradient is estimated from the difference in expectation values between these groups.

constructor

Initialize the branching simulator.

Arguments:

  • diff_method: method used for differentiating circuit parameters
  • num_samples: number of samples used to estimate expectation values
diff_methodattribute
diff_method: Literal['param_shift_inf', 'param_shift_single', 'param_shift_filter']
num_samplesattribute
num_samples: int
samplemethod
sample(
    circuit: CompiledBranchingPCircuit,
    x: Int[Array, 'pbits'],
    key: Key[Array, ''],
) -> Int[Array, 'num_samples num_pbits']

Obtain samples from the final distribution of the probabilistic circuit.

Arguments:

  • circuit: The probabilistic circuit to execute
  • x: The initial computational basis state of the circuit
  • key: The random key to use to obtain samples

Returns:

An integer array with shape (num_samples, num_pbits) containing the computational basis state samples.

expvalmethod
expval(
    circuit: CompiledBranchingPCircuit,
    x: Int[Array, 'pbits'],
    pbit: int,
    key: Key[Array, ''],
) -> Float[Array, '']

Estimate the expectation value of the given discrete site after circuit execution.

For a site with basis values $0, \ldots, d - 1$, this estimates $\mathbb{E}[x] = \sum_i i \, P(x = i)$.

Arguments:

  • circuit: The probabilistic circuit to execute
  • x: The initial computational basis state of the circuit
  • pbit: The index of the site to estimate the final expectation value of
  • key: The random key to use to obtain samples

Returns:

The expectation value of the given site after circuit execution.

expval_allmethod
expval_all(
    circuit: CompiledBranchingPCircuit,
    x: Int[Array, 'pbits'],
    key: Key[Array, ''],
) -> Float[Array, 'num_pbits']

Estimate the expectation value of all discrete sites after circuit execution.

For each site with basis values $0, \ldots, d - 1$, this estimates $\mathbb{E}[x] = \sum_i i \, P(x = i)$.

Arguments:

  • circuit: The probabilistic circuit to execute
  • x: The initial computational basis state of the circuit
  • key: The random key to use to obtain samples

Returns:

An array containing the expectation values of all discrete sites.

build_circuitmethod
build_circuit(
    circuit: DiscretePCircuit,
    thetas: list[Float[Array, '...']],
) -> CompiledBranchingPCircuit

Compile circuit for branch sampling.

Arguments:

  • circuit: The probabilistic circuit to compile.
  • thetas: Per-gate parameters aligned with circuit.gates.

Returns:

The compiled circuit.

CompiledStateVectorPCircuitclass
CompiledStateVectorPCircuit(
    gates: list[AbstractDiscreteGate],
    thetas: list[Float[Array, '...']],
    num_pdits: int,
    dims: tuple[int, ...],
    reps: int,
)

Compiled probabilistic circuit class for the state vector simulator.

gatesattribute
gates: list[AbstractDiscreteGate]
thetasattribute
thetas: list[jaxFloat[Array, '...']]
num_pditsattribute
num_pdits: int
dimsattribute
dims: tuple[int, ...]
repsattribute
reps: int
from_pcircuitclassmethod
from_pcircuit(circuit: DiscretePCircuit, thetas: list[Float[Array, '...']]) -> Self

Compile the given probabilistic circuit with the given parameters.

Arguments:

  • circuit: The probabilistic circuit to compile
  • thetas: Per-gate parameters aligned with circuit.gates

Returns:

The compiled circuit.

to_pcircuitmethod
to_pcircuit(structure: DiscretePCircuit) -> DiscretePCircuit

Return a DiscretePCircuit with this compiled circuit's gates and reps.

StateVectorSimulatorclass
StateVectorSimulator()

Simulator for exact state vectors representing probability distributions.

build_circuitmethod
build_circuit(
    circuit: DiscretePCircuit,
    thetas: list[Float[Array, '...']],
) -> CompiledStateVectorPCircuit

Build the circuit for this simulator with the given parameters.

Arguments:

  • circuit: The circuit to build
  • thetas: The parameters to bake into the compiled circuit.

Returns:

The built circuit.

apply_gatestaticmethod
apply_gate(
    state: Float[Array, 'dimensions'],
    gate: AbstractDiscreteGate,
    theta: Float[Array, '...'],
    dims: tuple[int, ...],
) -> Float[Array, 'dimensions']

Apply gate to the state StateVector and return the resulting state.

Arguments:

  • state: The state to apply the gate to
  • gate: The gate to apply
  • theta: The gate's parameters
  • dims: The dimensions of all sites in the circuit

Returns:

The state after applying the gate.

densitymethod
density(
    circuit: CompiledStateVectorPCircuit,
    x: Float[Array, 'dimensions'],
) -> Float[Array, 'dimensions']

Compute the final distribution over computational basis states.

Arguments:

  • circuit: The probabilistic circuit to compute the distribution of
  • x: The initial state vector, also the initial distribution

Returns:

The final distribution of the circuit.

expvalmethod
expval(
    circuit: CompiledStateVectorPCircuit,
    x: Float[Array, 'dimensions'],
    pbit: int,
) -> Float[Array, '']

Compute the expectation value of the given discrete site after circuit execution.

For a site with basis values $0, \ldots, d - 1$, this returns $\sum_i i \, P(x = i)$. For binary sites, this is equivalent to the probability of measuring 1.

Arguments:

  • circuit: The probabilistic circuit to execute
  • x: The initial state vector, also the initial distribution
  • pbit: The index of the site to compute the final expectation value of

Returns:

The expectation value of the given site after circuit execution.

expval_allmethod
expval_all(
    circuit: CompiledStateVectorPCircuit,
    x: Float[Array, 'dimensions'],
) -> Float[Array, 'pbits']

Compute the expectation value of all discrete sites after circuit execution.

For each site with basis values $0, \ldots, d - 1$, this returns $\sum_i i \, P(x = i)$. For binary sites, this is equivalent to the probability of measuring 1.

Arguments:

  • circuit: The probabilistic circuit to execute
  • x: The initial state vector, also the initial distribution

Returns:

An array containing the expectation values of all discrete sites.