use crate::quantum::{BaseGate, Circuit, Gate};
use std::f64::consts::PI;
fn hexf(v: f64) -> String {
format!("0x{:016X}", v.to_bits())
}
fn qubit(i: u8) -> String {
format!("%Qubit* inttoptr (i64 {i} to %Qubit*)")
}
fn p_angle(k: u16) -> f64 {
2.0 * PI / (1u64 << k.min(62)) as f64
}
struct Emitter {
body: Vec<String>,
used: Vec<&'static str>,
}
impl Emitter {
fn new() -> Self {
Emitter { body: Vec::new(), used: Vec::new() }
}
fn mark(&mut self, intrinsic: &'static str) {
if !self.used.contains(&intrinsic) {
self.used.push(intrinsic);
}
}
fn one(&mut self, name: &'static str, q: u8) {
self.mark(name);
self.body.push(format!(" call void @{name}({})", qubit(q)));
}
fn two(&mut self, name: &'static str, c: u8, t: u8) {
self.mark(name);
self.body.push(format!(" call void @{name}({}, {})", qubit(c), qubit(t)));
}
fn rz(&mut self, theta: f64, q: u8) {
self.mark("__quantum__qis__rz__body");
self.body.push(format!(
" call void @__quantum__qis__rz__body(double {}, {})",
hexf(theta),
qubit(q)
));
}
fn cy(&mut self, c: u8, t: u8) {
self.one("__quantum__qis__s__adj", t);
self.two("__quantum__qis__cnot__body", c, t);
self.one("__quantum__qis__s__body", t);
}
fn cp(&mut self, theta: f64, c: u8, t: u8) {
self.rz(theta / 2.0, c);
self.two("__quantum__qis__cnot__body", c, t);
self.rz(-theta / 2.0, t);
self.two("__quantum__qis__cnot__body", c, t);
self.rz(theta / 2.0, t);
}
fn ccx(&mut self, a: u8, b: u8, t: u8) {
const H: &str = "__quantum__qis__h__body";
const T: &str = "__quantum__qis__t__body";
const TD: &str = "__quantum__qis__t__adj";
const CN: &str = "__quantum__qis__cnot__body";
self.one(H, t);
self.two(CN, b, t);
self.one(TD, t);
self.two(CN, a, t);
self.one(T, t);
self.two(CN, b, t);
self.one(TD, t);
self.two(CN, a, t);
self.one(T, b);
self.one(T, t);
self.one(H, t);
self.two(CN, a, b);
self.one(T, a);
self.one(TD, b);
self.two(CN, a, b);
}
fn gate(&mut self, g: &Gate) {
let t = g.target;
match (g.base, g.controls.as_slice()) {
(BaseGate::I, _) => {}
(BaseGate::X, []) => self.one("__quantum__qis__x__body", t),
(BaseGate::Y, []) => self.one("__quantum__qis__y__body", t),
(BaseGate::Z, []) => self.one("__quantum__qis__z__body", t),
(BaseGate::H, []) => self.one("__quantum__qis__h__body", t),
(BaseGate::S, []) => self.one("__quantum__qis__s__body", t),
(BaseGate::Sdg, []) => self.one("__quantum__qis__s__adj", t),
(BaseGate::T, []) => self.one("__quantum__qis__t__body", t),
(BaseGate::Tdg, []) => self.one("__quantum__qis__t__adj", t),
(BaseGate::P, []) => self.rz(p_angle(g.param), t),
(BaseGate::X, [c]) => self.two("__quantum__qis__cnot__body", *c, t),
(BaseGate::Z, [c]) => self.two("__quantum__qis__cz__body", *c, t),
(BaseGate::Y, [c]) => self.cy(*c, t),
(BaseGate::P, [c]) => self.cp(p_angle(g.param), *c, t),
(BaseGate::X, [a, b]) => self.ccx(*a, *b, t),
(BaseGate::Z, [a, b]) => {
self.one("__quantum__qis__h__body", t);
self.ccx(*a, *b, t);
self.one("__quantum__qis__h__body", t);
}
(base, ctrls) => self.body.push(format!(
" ; UNMAPPED {base:?} controls={ctrls:?} target={t} — this module is NOT faithful"
)),
}
}
}
fn declaration(name: &str) -> String {
match name {
"__quantum__qis__rz__body" => format!("declare void @{name}(double, %Qubit*)"),
"__quantum__qis__cnot__body" | "__quantum__qis__cz__body" => {
format!("declare void @{name}(%Qubit*, %Qubit*)")
}
"__quantum__qis__mz__body" => format!("declare void @{name}(%Qubit*, %Result*)"),
"__quantum__rt__result_record_output" => format!("declare void @{name}(%Result*, i8*)"),
_ => format!("declare void @{name}(%Qubit*)"),
}
}
pub fn to_qir(c: &Circuit, measure_all: bool) -> String {
let mut e = Emitter::new();
for g in &c.ops {
e.gate(g);
}
let n = c.n_qubits;
if measure_all {
for q in 0..n {
e.mark("__quantum__qis__mz__body");
e.body.push(format!(
" call void @__quantum__qis__mz__body({}, %Result* inttoptr (i64 {q} to %Result*))",
qubit(q)
));
}
for q in 0..n {
e.mark("__quantum__rt__result_record_output");
e.body.push(format!(
" call void @__quantum__rt__result_record_output(%Result* inttoptr (i64 {q} to %Result*), i8* null)"
));
}
}
let mut s = String::new();
s.push_str("; QIR emitted by wai-quantum — deterministic, byte-identical on every target.\n");
s.push_str("source_filename = \"wai_quantum\"\n\n");
s.push_str("%Qubit = type opaque\n%Result = type opaque\n\n");
for i in &e.used {
s.push_str(&declaration(i));
s.push('\n');
}
s.push_str("\ndefine void @main() #0 {\nentry:\n");
for line in &e.body {
s.push_str(line);
s.push('\n');
}
s.push_str(" ret void\n}\n\n");
let results = if measure_all { n } else { 0 };
s.push_str(&format!(
"attributes #0 = {{ \"entry_point\" \"output_labeling_schema\" \
\"qir_profiles\"=\"base_profile\" \"required_num_qubits\"=\"{n}\" \
\"required_num_results\"=\"{results}\" }}\n\n"
));
s.push_str("!llvm.module.flags = !{!0, !1, !2, !3}\n");
s.push_str("!0 = !{i32 1, !\"qir_major_version\", i32 1}\n");
s.push_str("!1 = !{i32 7, !\"qir_minor_version\", i32 0}\n");
s.push_str("!2 = !{i32 1, !\"dynamic_qubit_management\", i1 false}\n");
s.push_str("!3 = !{i32 1, !\"dynamic_result_management\", i1 false}\n");
s
}
fn push(c: &mut Circuit, base: BaseGate, controls: Vec<u8>, target: u8, param: u16) {
c.ops.push(Gate { base, controls, target, param });
}
pub fn rewrite_cy(n: u8, c0: u8, t: u8) -> Circuit {
let mut c = Circuit::new(n);
push(&mut c, BaseGate::Sdg, vec![], t, 0);
c.cx(c0, t);
push(&mut c, BaseGate::S, vec![], t, 0);
c
}
pub fn rewrite_ccx(n: u8, a: u8, b: u8, t: u8) -> Circuit {
let mut c = Circuit::new(n);
c.h(t);
c.cx(b, t);
push(&mut c, BaseGate::Tdg, vec![], t, 0);
c.cx(a, t);
c.t(t);
c.cx(b, t);
push(&mut c, BaseGate::Tdg, vec![], t, 0);
c.cx(a, t);
c.t(b);
c.t(t);
c.h(t);
c.cx(a, b);
c.t(a);
push(&mut c, BaseGate::Tdg, vec![], b, 0);
c.cx(a, b);
c
}
#[cfg(test)]
mod tests {
use super::*;
use crate::quantum::ONE;
use crate::quantum_toolchain as qt;
fn agrees_on_all_inputs(n: u8, build: &dyn Fn(&mut Circuit), rewrite: &Circuit) -> (bool, i64) {
let tol = ONE / 100_000; let mut worst = ONE;
for k in 0..(1usize << n) {
let mut a = Circuit::new(n);
let mut b = Circuit::new(n);
for q in 0..n {
if k >> q & 1 == 1 {
a.x(q);
b.x(q);
}
}
build(&mut a);
for g in &rewrite.ops {
b.ops.push(g.clone());
}
let f = a.simulate().unwrap().fidelity_fx(&b.simulate().unwrap());
worst = worst.min(f);
if (ONE - f).abs() > tol {
return (false, f);
}
}
(true, worst)
}
#[test]
fn cy_rewrite_is_exact() {
let (ok, f) = agrees_on_all_inputs(
2,
&|c: &mut Circuit| {
c.ops.push(Gate { base: BaseGate::Y, controls: vec![0], target: 1, param: 0 });
},
&rewrite_cy(2, 0, 1),
);
assert!(ok, "CY rewrite fidelity {f} of {ONE}");
}
#[test]
fn ccx_rewrite_is_exact() {
let (ok, f) = agrees_on_all_inputs(3, &|c: &mut Circuit| { c.ccx(0, 1, 2); }, &rewrite_ccx(3, 0, 1, 2));
assert!(ok, "CCX rewrite fidelity {f} of {ONE}");
}
#[test]
fn ccx_rewrite_holds_with_controls_swapped() {
let (ok, f) = agrees_on_all_inputs(3, &|c: &mut Circuit| { c.ccx(1, 0, 2); }, &rewrite_ccx(3, 1, 0, 2));
assert!(ok, "CCX(swapped) rewrite fidelity {f} of {ONE}");
}
type C = (f64, f64);
fn cmul(a: C, b: C) -> C { (a.0 * b.0 - a.1 * b.1, a.0 * b.1 + a.1 * b.0) }
fn expi(t: f64) -> C { (t.cos(), t.sin()) }
fn rz_on(v: &mut [C; 4], phi: f64, q: usize) {
for (i, e) in v.iter_mut().enumerate() {
let bit = if q == 0 { i >> 1 & 1 } else { i & 1 };
let sign = if bit == 1 { 1.0 } else { -1.0 };
*e = cmul(*e, expi(sign * phi / 2.0));
}
}
fn cnot(v: &mut [C; 4]) {
v.swap(0b10, 0b11);
}
#[test]
fn cp_mapping_is_exact_up_to_global_phase() {
for &k in &[1u16, 2, 3, 5, 8] {
let theta = p_angle(k);
let mut global: Option<C> = None;
for basis in 0..4usize {
let mut v: [C; 4] = [(0.0, 0.0); 4];
v[basis] = (1.0, 0.0);
rz_on(&mut v, theta / 2.0, 0);
cnot(&mut v);
rz_on(&mut v, -theta / 2.0, 1);
cnot(&mut v);
rz_on(&mut v, theta / 2.0, 1);
let want: C = if basis == 3 { expi(theta) } else { (1.0, 0.0) };
let got = v[basis];
let ratio = cmul(got, (want.0, -want.1)); match global {
None => global = Some(ratio),
Some(g) => {
assert!((g.0 - ratio.0).abs() < 1e-12 && (g.1 - ratio.1).abs() < 1e-12,
"k={k} basis={basis}: phase {ratio:?} != {g:?} — not a GLOBAL phase");
}
}
for (j, e) in v.iter().enumerate() {
if j != basis {
assert!(e.0.abs() < 1e-12 && e.1.abs() < 1e-12, "k={k} leaked into {j}");
}
}
}
let g = global.unwrap();
assert!((g.0 * g.0 + g.1 * g.1 - 1.0).abs() < 1e-12, "global phase must be unit");
}
}
#[test]
fn emits_a_wellformed_module() {
let c = qt::ghz(3);
let ir = to_qir(&c, true);
assert!(ir.contains("%Qubit = type opaque"));
assert!(ir.contains("declare void @__quantum__qis__h__body(%Qubit*)"));
assert!(ir.contains("declare void @__quantum__qis__cnot__body(%Qubit*, %Qubit*)"));
assert!(ir.contains("define void @main() #0 {"));
assert!(ir.contains("\"required_num_qubits\"=\"3\""));
assert!(ir.contains("\"required_num_results\"=\"3\""));
assert!(ir.contains("qir_major_version"));
assert_eq!(ir.matches("__quantum__qis__cnot__body(%Qubit* inttoptr").count(), 2);
assert!(!ir.contains("UNMAPPED"));
}
#[test]
fn angles_are_exact_hex_floats() {
let mut c = Circuit::new(1);
c.p(3, 0); let ir = to_qir(&c, false);
let want = hexf(std::f64::consts::FRAC_PI_4);
assert!(ir.contains(&format!("double {want}")), "{ir}");
let bits = u64::from_str_radix(&want[2..], 16).unwrap();
assert_eq!(f64::from_bits(bits), std::f64::consts::FRAC_PI_4);
}
#[test]
fn every_toolchain_algorithm_maps_completely() {
for (id, _, _) in qt::catalog() {
if let Some(c) = qt::build_algorithm(id, 3, 5) {
let ir = to_qir(&c, true);
assert!(!ir.contains("UNMAPPED"), "{id} produced an unfaithful module:\n{ir}");
}
}
}
#[test]
fn no_symbol_is_declared_twice() {
for (id, _, _) in qt::catalog() {
if let Some(c) = qt::build_algorithm(id, 3, 5) {
for measure in [true, false] {
let ir = to_qir(&c, measure);
let mut names: Vec<&str> = ir
.lines()
.filter(|l| l.starts_with("declare "))
.map(|l| l.split('@').nth(1).unwrap().split('(').next().unwrap())
.collect();
let n = names.len();
names.sort_unstable();
names.dedup();
assert_eq!(n, names.len(), "{id}: duplicate declare in\n{ir}");
}
}
}
}
#[test]
fn deterministic() {
let c = qt::grover(3, 5);
assert_eq!(to_qir(&c, true), to_qir(&c, true));
}
}