prism-q 0.32.0

Fast Rust quantum circuit simulator. OpenQASM 3.0, multiple backends, AVX2 SIMD kernels, optional CUDA and MPI, QEC tooling, Python bindings.
Documentation
//! Controlled-phase batching passes: [`fuse_controlled_phases`] emits
//! [`Gate::BatchPhase`], [`batch_post_phase_1q`] re-batches trailing 1q runs.
//! The pass pipeline lives in [`crate::circuit::fusion`].

use std::borrow::Cow;

use num_complex::Complex64;

use super::{Circuit, Instruction, SmallVec, smallvec};
use crate::gates::{BatchPhaseData, Gate, MultiFusedData, is_diagonal_2x2};

use super::fusion::push_unique;
use super::plan::{Place, Tracer};

const MIN_BATCH_PHASES: usize = 2;

type PhaseVec = SmallVec<[(usize, Complex64); 8]>;
type TargetUserVec = SmallVec<[usize; 8]>;
type PendingPhaseVec = Vec<Option<PhaseVec>>;

fn is_controlled_phase_2q(inst: &Instruction) -> bool {
    matches!(
        inst,
        Instruction::Gate { gate, .. }
            if gate.controlled_phase().is_some() && gate.num_qubits() == 2
    )
}

fn remove_target_user(target_users: &mut [TargetUserVec], target: usize, control: usize) {
    target_users[target].retain(|c| *c != control);
}

fn emit_phase_chain(
    control: usize,
    phases: PhaseVec,
    output: &mut Vec<Instruction>,
    changed: &mut bool,
) {
    // Every entry is diagonal on the same control, so a chain longer than the
    // kernel group tables hold splits into consecutive batches.
    if phases.len() > BatchPhaseData::MAX_PHASES {
        for chunk in phases.chunks(BatchPhaseData::MAX_PHASES) {
            emit_phase_batch(control, PhaseVec::from_slice(chunk), output, changed);
        }
        return;
    }
    emit_phase_batch(control, phases, output, changed);
}

fn emit_phase_batch(
    control: usize,
    phases: PhaseVec,
    output: &mut Vec<Instruction>,
    changed: &mut bool,
) {
    if phases.len() >= MIN_BATCH_PHASES {
        output.push(Instruction::Gate {
            gate: Gate::BatchPhase(Box::new(BatchPhaseData { phases })),
            targets: smallvec![control],
        });
        *changed = true;
        return;
    }

    let one = Complex64::new(1.0, 0.0);
    let zero = Complex64::new(0.0, 0.0);
    for (target, phase) in phases {
        output.push(Instruction::Gate {
            gate: Gate::cu([[one, zero], [zero, phase]]),
            targets: smallvec![control, target],
        });
    }
}

fn flush_phase_control(
    control: usize,
    pending: &mut [Option<PhaseVec>],
    target_users: &mut [TargetUserVec],
    output: &mut Vec<Instruction>,
    changed: &mut bool,
) {
    let Some(phases) = pending[control].take() else {
        return;
    };
    for &(target, _) in &phases {
        remove_target_user(target_users, target, control);
    }
    emit_phase_chain(control, phases, output, changed);
}

fn push_pending_phase(
    control: usize,
    target: usize,
    phase: Complex64,
    pending: &mut [Option<PhaseVec>],
    target_users: &mut [TargetUserVec],
) {
    // The kernel indexes one bit per distinct target, so a repeated pair has to
    // fold into the entry already there rather than push a second one.
    match &mut pending[control] {
        Some(v) => {
            if let Some(entry) = v.iter_mut().find(|(t, _)| *t == target) {
                entry.1 *= phase;
                return;
            }
            v.push((target, phase));
        }
        slot => *slot = Some(smallvec![(target, phase)]),
    }
    target_users[target].push(control);
}

fn flush_phase_target_conflicts(
    target: usize,
    pending: &mut [Option<PhaseVec>],
    target_users: &mut [TargetUserVec],
    output: &mut Vec<Instruction>,
    changed: &mut bool,
) {
    let controls = std::mem::take(&mut target_users[target]);
    if controls.is_empty() {
        return;
    }

    let mut re_rooted: PhaseVec = SmallVec::new();
    for control in controls {
        let Some(phases) = pending[control].take() else {
            continue;
        };
        let mut kept: PhaseVec = SmallVec::new();
        for (phase_target, phase) in phases {
            if phase_target == target {
                re_rooted.push((control, phase));
            } else {
                kept.push((phase_target, phase));
            }
        }
        if !kept.is_empty() {
            pending[control] = Some(kept);
        }
    }

    if !re_rooted.is_empty() {
        emit_phase_chain(target, re_rooted, output, changed);
    }
}

fn flush_phase_qubits_in_use(
    qs: &[usize],
    diagonal_only: bool,
    pending: &mut [Option<PhaseVec>],
    target_users: &mut [TargetUserVec],
    output: &mut Vec<Instruction>,
    changed: &mut bool,
) {
    for &q in qs {
        flush_phase_control(q, pending, target_users, output, changed);
    }
    if diagonal_only {
        return;
    }
    for &q in qs {
        flush_phase_target_conflicts(q, pending, target_users, output, changed);
    }
}

/// Returns the input unchanged unless at least one `BatchPhase` is emitted.
pub fn fuse_controlled_phases<'a>(circuit: Cow<'a, Circuit>, t: &mut Tracer) -> Cow<'a, Circuit> {
    if !circuit.instructions.iter().any(is_controlled_phase_2q) {
        return circuit;
    }

    let mut output: Vec<Instruction> = Vec::with_capacity(circuit.instructions.len());
    let n = circuit.num_qubits;
    let mut pending: PendingPhaseVec = (0..n).map(|_| None).collect();
    let mut target_users: Vec<TargetUserVec> = (0..n).map(|_| SmallVec::new()).collect();
    let mut changed = false;

    for inst in &circuit.instructions {
        match inst {
            Instruction::Gate { gate, targets } => {
                if let Some(phase) = gate.controlled_phase() {
                    if gate.num_qubits() == 2 {
                        let control = targets[0];
                        let target = targets[1];
                        // Incoming cphase on target, its action depends on
                        // the current diagonal-frame phase of `target`, so
                        // any pending control chain on `target` must flush first.
                        flush_phase_qubits_in_use(
                            std::slice::from_ref(&target),
                            true,
                            &mut pending,
                            &mut target_users,
                            &mut output,
                            &mut changed,
                        );
                        push_pending_phase(control, target, phase, &mut pending, &mut target_users);
                        continue;
                    }
                }
                let diagonal_only = gate.num_qubits() == 1 && gate.is_diagonal_1q();
                flush_phase_qubits_in_use(
                    targets,
                    diagonal_only,
                    &mut pending,
                    &mut target_users,
                    &mut output,
                    &mut changed,
                );
                output.push(inst.clone());
            }
            Instruction::Measure { qubit, .. } | Instruction::Reset { qubit } => {
                flush_phase_qubits_in_use(
                    std::slice::from_ref(qubit),
                    false,
                    &mut pending,
                    &mut target_users,
                    &mut output,
                    &mut changed,
                );
                output.push(inst.clone());
            }
            Instruction::Barrier { qubits } => {
                flush_phase_qubits_in_use(
                    qubits,
                    false,
                    &mut pending,
                    &mut target_users,
                    &mut output,
                    &mut changed,
                );
                output.push(inst.clone());
            }
            Instruction::Region(region) => {
                flush_phase_qubits_in_use(
                    region.qubits(),
                    false,
                    &mut pending,
                    &mut target_users,
                    &mut output,
                    &mut changed,
                );
                output.push(inst.clone());
            }
            Instruction::Conditional { targets, .. } => {
                flush_phase_qubits_in_use(
                    targets,
                    false,
                    &mut pending,
                    &mut target_users,
                    &mut output,
                    &mut changed,
                );
                output.push(inst.clone());
            }
        }
    }

    for q in 0..n {
        flush_phase_control(
            q,
            &mut pending,
            &mut target_users,
            &mut output,
            &mut changed,
        );
    }

    if changed {
        t.bail();
        Cow::Owned(circuit.with_instructions(output))
    } else {
        circuit
    }
}

fn is_batchable_1q(inst: &Instruction) -> bool {
    matches!(
        inst,
        Instruction::Gate { targets, gate, .. }
            if targets.len() == 1
                && !matches!(gate, Gate::BatchPhase(_) | Gate::MultiFused(_) | Gate::Multi2q(_) | Gate::DiagonalBatch(_))
    )
}

pub(super) fn batch_post_phase_1q<'a>(
    circuit: Cow<'a, Circuit>,
    t: &mut Tracer,
) -> Cow<'a, Circuit> {
    let mut max_run = 0usize;
    let mut run = 0usize;
    for inst in &circuit.instructions {
        if is_batchable_1q(inst) {
            run += 1;
            max_run = max_run.max(run);
        } else {
            run = 0;
        }
    }
    if max_run < 2 {
        return circuit;
    }

    let mut output: Vec<Instruction> = Vec::with_capacity(circuit.instructions.len());
    let mut pending: Vec<(usize, [[Complex64; 2]; 2])> = Vec::new();
    let mut pending_src: Vec<usize> = Vec::new();
    t.begin();

    let flush = |pending: &mut Vec<(usize, [[Complex64; 2]; 2])>,
                 pending_src: &mut Vec<usize>,
                 output: &mut Vec<Instruction>,
                 t: &mut Tracer| {
        if pending.len() >= 2 {
            let mut targets: SmallVec<[usize; 4]> = SmallVec::new();
            for &(q, _) in pending.iter() {
                push_unique(&mut targets, q);
            }
            targets.sort_unstable();
            let all_diagonal = pending.iter().all(|(_, m)| is_diagonal_2x2(m));
            output.push(Instruction::Gate {
                gate: Gate::MultiFused(Box::new(MultiFusedData {
                    gates: std::mem::take(pending),
                    all_diagonal,
                })),
                targets,
            });
            if t.on {
                let entries: Vec<Vec<(usize, Place)>> = pending_src
                    .drain(..)
                    .map(|src| vec![(src, Place::Plain)])
                    .collect();
                t.batch(&entries);
            }
        } else {
            for (k, (q, mat)) in pending.drain(..).enumerate() {
                output.push(Instruction::Gate {
                    gate: Gate::Fused(Box::new(mat)),
                    targets: smallvec![q],
                });
                if t.on {
                    t.keep(pending_src[k]);
                }
            }
            pending_src.clear();
        }
    };

    for (i, inst) in circuit.instructions.iter().enumerate() {
        if is_batchable_1q(inst) {
            if let Instruction::Gate { gate, targets, .. } = inst {
                pending.push((targets[0], gate.matrix_2x2()));
                t.note_idx(&mut pending_src, i);
            }
        } else {
            flush(&mut pending, &mut pending_src, &mut output, t);
            output.push(inst.clone());
            t.keep(i);
        }
    }

    flush(&mut pending, &mut pending_src, &mut output, t);
    t.commit();

    Cow::Owned(circuit.with_instructions(output))
}