use std::collections::{BTreeMap, HashMap, HashSet};
use crate::errors::{Result, TicitError};
use crate::factored::{FactoredInstruction, FactoredInstructionProgram};
use crate::symbolic::{SymbolicBool, SymbolicBoolEvaluationPlan, symbolic_bool, xor_bool};
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct MeasurementParity {
pub records: Vec<usize>,
pub value: bool,
}
impl MeasurementParity {
#[must_use]
pub fn new(records: impl Into<Vec<usize>>, value: bool) -> Self {
Self {
records: records.into(),
value,
}
}
}
#[derive(Clone, Debug, Default)]
pub(crate) struct ForcedBranch {
pub instruction: usize,
pub plan: SymbolicBoolEvaluationPlan,
}
#[derive(Clone, Copy, Debug)]
enum SymbolSource {
Branch,
Derived(usize),
}
struct ProgramSymbols {
sources: HashMap<i32, SymbolSource>,
branch_instruction: HashMap<i32, usize>,
record_instruction: HashMap<i32, usize>,
expansions: HashMap<i32, SymbolicBool>,
}
impl ProgramSymbols {
fn new(program: &FactoredInstructionProgram) -> Self {
let mut sources = HashMap::new();
let mut branch_instruction = HashMap::new();
let mut record_instruction = HashMap::new();
for (index, instruction) in program.instructions.iter().enumerate() {
let branch = match instruction {
FactoredInstruction::MeasurePrecomputedActivePauli(inst) => Some(inst.branch),
FactoredInstruction::IntroduceDormantMeasurementBranch(inst) => Some(inst.branch),
_ => None,
};
let probes = instruction.exp_val().is_some();
if let Some(branch) = branch
&& !probes
{
sources.insert(branch, SymbolSource::Branch);
branch_instruction.insert(branch, index);
}
if let Some(condition) = instruction.record_condition()
&& !probes
{
sources
.entry(condition)
.or_insert(SymbolSource::Derived(index));
}
if let Some(record) = instruction.record()
&& !probes
{
record_instruction.insert(record, index);
}
}
Self {
sources,
branch_instruction,
record_instruction,
expansions: HashMap::new(),
}
}
fn record_expansion(
&mut self,
program: &FactoredInstructionProgram,
record: usize,
) -> Result<SymbolicBool> {
let one_based = i32::try_from(record + 1)
.map_err(|_| TicitError::new("pinned measurement record index is out of range"))?;
let Some(&instruction) = self.record_instruction.get(&one_based) else {
return Err(TicitError::new(format!(
"pinned measurement record {record} is not written by this circuit"
)));
};
let outcome = program.instructions[instruction]
.outcome()
.ok_or_else(|| TicitError::new("measurement instruction has no outcome expression"))?
.clone();
self.expand(program, &outcome)
}
fn expand(
&mut self,
program: &FactoredInstructionProgram,
expr: &SymbolicBool,
) -> Result<SymbolicBool> {
for &condition in &expr.conditions {
self.expand_symbol(program, condition)?;
}
let mut out = SymbolicBool::from(expr.constant);
for &condition in &expr.conditions {
let expansion = self
.expansions
.get(&condition)
.expect("every condition was just expanded");
out = xor_bool(&out, expansion);
}
Ok(out)
}
fn expand_symbol(&mut self, program: &FactoredInstructionProgram, symbol: i32) -> Result<()> {
let mut stack = vec![symbol];
let mut in_progress: HashSet<i32> = HashSet::new();
while let Some(&top) = stack.last() {
if self.expansions.contains_key(&top) {
in_progress.remove(&top);
stack.pop();
continue;
}
match self.sources.get(&top).copied() {
None => {
self.expansions.insert(top, SymbolicBool::from(false));
}
Some(SymbolSource::Branch) => {
self.expansions.insert(top, symbolic_bool(top));
}
Some(SymbolSource::Derived(instruction)) => {
in_progress.insert(top);
let outcome = program.instructions[instruction].outcome().ok_or_else(|| {
TicitError::new("record condition has no outcome expression")
})?;
let mut pending = 0usize;
for &condition in &outcome.conditions {
if self.expansions.contains_key(&condition) {
continue;
}
if in_progress.contains(&condition) {
return Err(TicitError::new("measurement record expression is cyclic"));
}
stack.push(condition);
pending += 1;
}
if pending > 0 {
continue;
}
let mut value = SymbolicBool::from(outcome.constant);
for &condition in &outcome.conditions {
let expansion = self
.expansions
.get(&condition)
.expect("all conditions resolved above");
value = xor_bool(&value, expansion);
}
self.expansions.insert(top, value);
}
}
in_progress.remove(&top);
stack.pop();
}
Ok(())
}
fn last_drawn(&self, conditions: &[i32]) -> Option<(usize, i32)> {
conditions
.iter()
.filter_map(|&symbol| {
self.branch_instruction
.get(&symbol)
.map(|&instruction| (instruction, symbol))
})
.max()
}
}
pub(crate) fn plan_pinned_measurements(
program: &FactoredInstructionProgram,
constraints: &[MeasurementParity],
) -> Result<Vec<ForcedBranch>> {
if constraints.is_empty() {
return Ok(Vec::new());
}
let mut symbols = ProgramSymbols::new(program);
let mut pivot_rows: BTreeMap<usize, (i32, SymbolicBool)> = BTreeMap::new();
for constraint in constraints {
let mut row = SymbolicBool::from(constraint.value);
for &record in &constraint.records {
let expansion = symbols.record_expansion(program, record)?;
row = xor_bool(&row, &expansion);
}
loop {
let Some((instruction, pivot)) = symbols.last_drawn(&row.conditions) else {
if row.constant {
return Err(TicitError::new(format!(
"pinned measurement parity over {:?} is deterministic and cannot be {}",
constraint.records,
u8::from(constraint.value),
)));
}
break;
};
match pivot_rows.get(&instruction) {
Some((_, existing)) => row = xor_bool(&row, existing),
None => {
pivot_rows.insert(instruction, (pivot, row));
break;
}
}
}
}
let mut forced = Vec::with_capacity(pivot_rows.len());
for (instruction, (pivot, row)) in pivot_rows {
let assignment = xor_bool(&row, &symbolic_bool(pivot));
forced.push(ForcedBranch {
instruction,
plan: SymbolicBoolEvaluationPlan::new(&assignment),
});
}
Ok(forced)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::circuit::{parse_ticit_text, plan_ticit_factored_program};
fn planned(text: &str) -> FactoredInstructionProgram {
let parsed = parse_ticit_text(text).expect("test circuit parses");
plan_ticit_factored_program(&parsed).expect("test circuit plans")
}
#[test]
fn a_free_coin_is_pinned_at_its_own_instruction() {
let program = planned("H 0\nM 0\n");
let forced = plan_pinned_measurements(&program, &[MeasurementParity::new([0], true)])
.expect("a fair coin can be pinned");
assert_eq!(forced.len(), 1);
assert!(forced[0].plan.conditions.is_empty());
assert!(forced[0].plan.constant);
}
#[test]
fn a_deterministic_record_rejects_the_wrong_value() {
let program = planned("M 0\n");
let error = plan_pinned_measurements(&program, &[MeasurementParity::new([0], true)])
.expect_err("|0> always measures 0");
assert!(error.to_string().contains("deterministic"));
}
#[test]
fn a_deterministic_record_accepts_the_right_value() {
let program = planned("M 0\n");
let forced = plan_pinned_measurements(&program, &[MeasurementParity::new([0], false)])
.expect("|0> always measures 0");
assert!(forced.is_empty());
}
#[test]
fn a_parity_pins_only_its_last_free_coin() {
let program = planned("H 0\nH 1\nM 0\nM 1\n");
let forced = plan_pinned_measurements(&program, &[MeasurementParity::new([0, 1], true)])
.expect("two fair coins can meet a parity");
assert_eq!(forced.len(), 1, "only the last draw is pinned");
assert_eq!(forced[0].plan.conditions.len(), 1, "it follows the first");
}
#[test]
fn independent_constraints_take_distinct_pivots() {
let program = planned("H 0\nH 1\nM 0\nM 1\n");
let forced = plan_pinned_measurements(
&program,
&[
MeasurementParity::new([0], true),
MeasurementParity::new([0, 1], false),
],
)
.expect("independent constraints solve");
assert_eq!(forced.len(), 2);
assert_ne!(forced[0].instruction, forced[1].instruction);
}
#[test]
fn contradictory_constraints_are_rejected() {
let program = planned("H 0\nM 0\nM 0\n");
let error = plan_pinned_measurements(
&program,
&[
MeasurementParity::new([0], true),
MeasurementParity::new([1], false),
],
)
.expect_err("the second measurement repeats the first");
assert!(error.to_string().contains("deterministic"));
}
#[test]
fn an_unwritten_record_is_rejected() {
let program = planned("H 0\nM 0\n");
let error = plan_pinned_measurements(&program, &[MeasurementParity::new([7], true)])
.expect_err("record 7 does not exist");
assert!(error.to_string().contains("not written"));
}
}