pub use crate::linalg::C;
use crate::repro::ln;
use std::collections::BTreeMap;
pub use crate::gates::{complex, Gate};
#[derive(Clone, Debug)]
pub struct Tensor {
pub inds: Vec<u32>,
pub data: Vec<C>,
}
#[derive(Clone, Debug)]
pub struct Network {
pub tensors: Vec<Tensor>,
pub dims: Vec<usize>,
pub open: Vec<u32>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum TnError {
BadGate(usize),
BadOutput,
PlanMismatch,
}
impl Network {
pub fn amplitude(n: u32, gates: &[Gate], bits: &[u8], open_qubits: &[u32]) -> Result<Network, TnError> {
if bits.len() != n as usize || bits.iter().any(|&b| b > 1) || open_qubits.iter().any(|&q| q >= n) {
return Err(TnError::BadOutput);
}
let mut next = 0u32;
let mut fresh = || {
next += 1;
next - 1
};
let mut tensors = Vec::new();
let mut wire: Vec<u32> = (0..n).map(|_| fresh()).collect();
for &w in &wire {
tensors.push(Tensor { inds: vec![w], data: vec![C::ONE, C::ZERO] });
}
for (gi, g) in gates.iter().enumerate() {
let k = g.qubits.len();
let mut seen = g.qubits.clone();
seen.sort_unstable();
seen.dedup();
if k == 0 || seen.len() != k || g.qubits.iter().any(|&q| q >= n) || g.matrix.len() != 1 << (2 * k) {
return Err(TnError::BadGate(gi));
}
let outs: Vec<u32> = (0..k).map(|_| fresh()).collect();
let mut inds = outs.clone();
inds.extend(g.qubits.iter().map(|&q| wire[q as usize]));
tensors.push(Tensor { inds, data: g.matrix.clone() });
for (i, &q) in g.qubits.iter().enumerate() {
wire[q as usize] = outs[i];
}
}
let mut open = Vec::new();
for q in 0..n {
if !open_qubits.contains(&q) {
let v = if bits[q as usize] == 0 { vec![C::ONE, C::ZERO] } else { vec![C::ZERO, C::ONE] };
tensors.push(Tensor { inds: vec![wire[q as usize]], data: v });
}
}
for &q in open_qubits {
open.push(wire[q as usize]);
}
Ok(Network { tensors, dims: vec![2; next as usize], open })
}
pub fn simplify(&mut self) {
loop {
let owners = self.owners();
let mut done = true;
for t in 0..self.tensors.len() {
let rank = self.tensors[t].inds.len();
if rank > 2 || self.tensors.len() == 1 {
continue;
}
let mut target = None;
for &i in &self.tensors[t].inds {
if let Some(&u) = owners.get(&i).and_then(|o| o.iter().find(|&&u| u != t)) {
let merged = sym_diff(&self.tensors[t].inds, &self.tensors[u].inds, &self.open).len();
if merged <= self.tensors[u].inds.len() {
target = Some(u);
break;
}
}
}
if let Some(u) = target {
let (a, b) = if t < u { (t, u) } else { (u, t) };
let tb = self.tensors.remove(b);
let ta = self.tensors.remove(a);
let merged = if ta.inds.len() >= tb.inds.len() {
contract(&ta, &tb, &self.dims, &self.open)
} else {
contract(&tb, &ta, &self.dims, &self.open)
};
self.tensors.insert(a, merged);
done = false;
break;
}
}
if done {
return;
}
}
}
fn owners(&self) -> BTreeMap<u32, Vec<usize>> {
let mut o: BTreeMap<u32, Vec<usize>> = BTreeMap::new();
for (t, x) in self.tensors.iter().enumerate() {
for &i in &x.inds {
o.entry(i).or_default().push(t);
}
}
o
}
pub fn shape(&self) -> Vec<Vec<u32>> {
self.tensors.iter().map(|t| t.inds.clone()).collect()
}
}
fn sym_diff(a: &[u32], b: &[u32], open: &[u32]) -> Vec<u32> {
let mut out: Vec<u32> = a.iter().copied().filter(|i| !b.contains(i) || open.contains(i)).collect();
out.extend(b.iter().copied().filter(|i| !a.contains(i)));
out
}
fn permute(t: &Tensor, order: &[u32], dims: &[usize]) -> Vec<C> {
if t.inds == order {
return t.data.clone();
}
let r = t.inds.len();
let mut src_stride = vec![1usize; r];
for p in (0..r.saturating_sub(1)).rev() {
src_stride[p] = src_stride[p + 1] * dims[t.inds[p + 1] as usize];
}
let pos: Vec<usize> = order.iter().map(|i| t.inds.iter().position(|j| j == i).unwrap()).collect();
let out_dims: Vec<usize> = order.iter().map(|&i| dims[i as usize]).collect();
let total = t.data.len();
let mut out = Vec::with_capacity(total);
let mut idx = vec![0usize; r];
let mut src = 0usize;
for _ in 0..total {
out.push(t.data[src]);
for p in (0..r).rev() {
idx[p] += 1;
src += src_stride[pos[p]];
if idx[p] < out_dims[p] {
break;
}
src -= src_stride[pos[p]] * out_dims[p];
idx[p] = 0;
}
}
out
}
pub fn contract(a: &Tensor, b: &Tensor, dims: &[usize], open: &[u32]) -> Tensor {
let shared: Vec<u32> = a.inds.iter().copied().filter(|i| b.inds.contains(i) && !open.contains(i)).collect();
let a_rem: Vec<u32> = a.inds.iter().copied().filter(|i| !shared.contains(i)).collect();
let b_rem: Vec<u32> = b.inds.iter().copied().filter(|i| !shared.contains(i)).collect();
let size = |v: &[u32]| v.iter().map(|&i| dims[i as usize]).product::<usize>();
let (m, k, n) = (size(&a_rem), size(&shared), size(&b_rem));
let mut ao = a_rem.clone();
ao.extend(&shared);
let mut bo = shared.clone();
bo.extend(&b_rem);
let ad = permute(a, &ao, dims);
let bd = permute(b, &bo, dims);
let mut out = vec![C::ZERO; m * n];
for i in 0..m {
let row = &mut out[i * n..(i + 1) * n];
for kk in 0..k {
let x = ad[i * k + kk];
if x.re == 0.0 && x.im == 0.0 {
continue;
}
let brow = &bd[kk * n..(kk + 1) * n];
for (o, y) in row.iter_mut().zip(brow) {
o.re += x.re * y.re - x.im * y.im;
o.im += x.re * y.im + x.im * y.re;
}
}
}
let mut inds = a_rem;
inds.extend(b_rem);
Tensor { inds, data: out }
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct Path {
pub steps: Vec<(usize, usize)>,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Cost {
pub flops: f64,
pub max_size: f64,
}
impl Cost {
pub fn log10_flops(&self) -> f64 {
ln(self.flops) / core::f64::consts::LN_10
}
pub fn log2_size(&self) -> f64 {
ln(self.max_size) / core::f64::consts::LN_2
}
}
fn set_size(s: &[u32], dims: &[usize]) -> f64 {
s.iter().map(|&i| dims[i as usize] as f64).product()
}
fn union(a: &[u32], b: &[u32]) -> Vec<u32> {
let mut u = a.to_vec();
u.extend(b.iter().copied().filter(|i| !a.contains(i)));
u
}
pub fn path_cost(shape: &[Vec<u32>], dims: &[usize], open: &[u32], path: &Path) -> Cost {
let mut live: Vec<Vec<u32>> = shape.to_vec();
let mut flops = 0.0;
let mut max_size: f64 = shape.iter().map(|s| set_size(s, dims)).fold(0.0, f64::max);
for &(a, b) in &path.steps {
flops += set_size(&union(&live[a], &live[b]), dims);
let out = sym_diff(&live[a], &live[b], open);
max_size = max_size.max(set_size(&out, dims));
live.push(out);
}
Cost { flops, max_size }
}
pub(crate) struct Rng(u64);
impl Rng {
pub(crate) fn new(seed: u64) -> Rng {
Rng(seed)
}
fn next(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9e37_79b9_7f4a_7c15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xbf58_476d_1ce4_e5b9);
z = (z ^ (z >> 27)).wrapping_mul(0x94d0_49bb_1331_11eb);
z ^ (z >> 31)
}
fn unit(&mut self) -> f64 {
((self.next() >> 11) as f64 + 0.5) * (1.0 / 9_007_199_254_740_992.0)
}
fn below(&mut self, n: usize) -> usize {
(self.next() % n as u64) as usize
}
}
pub fn greedy_path(shape: &[Vec<u32>], dims: &[usize], open: &[u32], costmod: f64, temperature: f64, seed: u64) -> Path {
let mut rng = Rng::new(seed);
let mut live: Vec<Option<Vec<u32>>> = shape.iter().cloned().map(Some).collect();
let mut steps = Vec::new();
let ids: Vec<usize> = (0..shape.len()).collect();
greedy_steps(&ids, &mut live, &mut steps, dims, open, costmod, temperature, &mut rng);
Path { steps }
}
#[allow(clippy::too_many_arguments)]
fn greedy_steps(
ids: &[usize],
live: &mut Vec<Option<Vec<u32>>>,
steps: &mut Vec<(usize, usize)>,
dims: &[usize],
open: &[u32],
costmod: f64,
temperature: f64,
rng: &mut Rng,
) -> usize {
if ids.len() == 1 {
return ids[0];
}
let mut owners: BTreeMap<u32, Vec<usize>> = BTreeMap::new();
for &t in ids {
for &i in live[t].as_ref().unwrap() {
owners.entry(i).or_default().push(t);
}
}
let score = |a: &[u32], b: &[u32], rng: &mut Rng| -> f64 {
let out = set_size(&sym_diff(a, b, open), dims);
let c = out / costmod - costmod * (set_size(a, dims) + set_size(b, dims));
let c = if c >= 0.0 { ln(1.0 + c) } else { -ln(1.0 - c) };
if temperature > 0.0 { c - temperature * -ln(-ln(rng.unit())) } else { c }
};
let mut cands: Vec<(f64, usize, usize)> = Vec::new();
let mut seen = std::collections::BTreeSet::new();
for ts in owners.values() {
for x in 0..ts.len() {
for y in x + 1..ts.len() {
let (a, b) = (ts[x].min(ts[y]), ts[x].max(ts[y]));
if seen.insert((a, b)) {
let sc = score(live[a].as_ref().unwrap(), live[b].as_ref().unwrap(), rng);
cands.push((sc, a, b));
}
}
}
}
let mut members: Vec<usize> = ids.to_vec();
loop {
let mut best: Option<usize> = None;
for (ci, c) in cands.iter().enumerate() {
if best.is_none_or(|bi| {
let b = cands[bi];
c.0 < b.0 || (c.0 == b.0 && (c.1, c.2) < (b.1, b.2))
}) {
best = Some(ci);
}
}
let Some(bi) = best else { break };
let (_, a, b) = cands.swap_remove(bi);
cands.retain(|c| c.1 != a && c.1 != b && c.2 != a && c.2 != b);
let sa = live[a].take().unwrap();
let sb = live[b].take().unwrap();
let out = sym_diff(&sa, &sb, open);
let id = live.len();
steps.push((a, b));
members.retain(|&t| t != a && t != b);
for &i in &out {
if let Some(o) = owners.get_mut(&i) {
o.retain(|&t| t != a && t != b);
for &t in o.iter() {
let sc = score(live[t].as_ref().unwrap(), &out, rng);
cands.push((sc, t.min(id), t.max(id)));
}
o.push(id);
}
}
live.push(Some(out));
members.push(id);
}
while members.len() > 1 {
members.sort_by(|&x, &y| {
set_size(live[x].as_ref().unwrap(), dims)
.partial_cmp(&set_size(live[y].as_ref().unwrap(), dims))
.unwrap()
.then(x.cmp(&y))
});
let (a, b) = (members[0], members[1]);
let sa = live[a].take().unwrap();
let sb = live[b].take().unwrap();
steps.push((a, b));
let id = live.len();
live.push(Some(sym_diff(&sa, &sb, open)));
members.drain(0..2);
members.push(id);
}
members[0]
}
#[derive(Clone, Copy, Debug, PartialEq)]
struct PartitionParams {
imbalance: f64,
cutoff: usize,
passes: usize,
costmod: f64,
temperature: f64,
}
fn partition_path(shape: &[Vec<u32>], dims: &[usize], open: &[u32], p: PartitionParams, seed: u64) -> Path {
let mut rng = Rng::new(seed);
let mut live: Vec<Option<Vec<u32>>> = shape.iter().cloned().map(Some).collect();
let mut steps = Vec::new();
let ids: Vec<usize> = (0..shape.len()).collect();
partition_rec(&ids, &mut live, &mut steps, dims, open, &p, &mut rng);
Path { steps }
}
fn partition_rec(
ids: &[usize],
live: &mut Vec<Option<Vec<u32>>>,
steps: &mut Vec<(usize, usize)>,
dims: &[usize],
open: &[u32],
p: &PartitionParams,
rng: &mut Rng,
) -> usize {
if ids.len() <= p.cutoff.max(2) {
return greedy_steps(ids, live, steps, dims, open, p.costmod, p.temperature, rng);
}
let (left, right) = bisect(ids, live, dims, p, rng);
let a = partition_rec(&left, live, steps, dims, open, p, rng);
let b = partition_rec(&right, live, steps, dims, open, p, rng);
let sa = live[a].take().unwrap();
let sb = live[b].take().unwrap();
steps.push((a, b));
live.push(Some(sym_diff(&sa, &sb, open)));
live.len() - 1
}
struct WGraph {
nw: Vec<f64>,
adj: Vec<Vec<(usize, f64)>>,
}
impl WGraph {
fn n(&self) -> usize {
self.nw.len()
}
}
fn merged(adj: Vec<Vec<(usize, f64)>>) -> Vec<Vec<(usize, f64)>> {
adj.into_iter()
.map(|mut v| {
v.sort_by_key(|a| a.0);
let mut out: Vec<(usize, f64)> = Vec::with_capacity(v.len());
for (u, w) in v {
match out.last_mut() {
Some(last) if last.0 == u => last.1 += w,
_ => out.push((u, w)),
}
}
out
})
.collect()
}
fn coarsen(g: &WGraph, rng: &mut Rng) -> (WGraph, Vec<usize>) {
let n = g.n();
let mut order: Vec<usize> = (0..n).collect();
for i in (1..n).rev() {
order.swap(i, rng.below(i + 1));
}
let mut mate = vec![usize::MAX; n];
for &v in &order {
if mate[v] != usize::MAX {
continue;
}
let mut best: Option<(f64, usize)> = None;
for &(u, w) in &g.adj[v] {
if u != v && mate[u] == usize::MAX && best.is_none_or(|(bw, _)| w > bw) {
best = Some((w, u));
}
}
match best {
Some((_, u)) => {
mate[v] = u;
mate[u] = v;
}
None => mate[v] = v,
}
}
let mut map = vec![usize::MAX; n];
let mut nw = Vec::new();
for v in 0..n {
if map[v] == usize::MAX {
let c = nw.len();
map[v] = c;
let u = mate[v];
let mut w = g.nw[v];
if u != v {
map[u] = c;
w += g.nw[u];
}
nw.push(w);
}
}
let mut adj = vec![Vec::new(); nw.len()];
for v in 0..n {
for &(u, w) in &g.adj[v] {
if map[v] != map[u] {
adj[map[v]].push((map[u], w));
}
}
}
(WGraph { nw, adj: merged(adj) }, map)
}
fn refine(g: &WGraph, side: &mut Vec<bool>, lo: f64, hi: f64, passes: usize) {
let n = g.n();
let cut = |s: &[bool]| -> f64 {
let mut c = 0.0;
for v in 0..n {
for &(u, w) in &g.adj[v] {
if v < u && s[v] != s[u] {
c += w;
}
}
}
c
};
for _ in 0..passes {
let mut locked = vec![false; n];
let mut cur = side.clone();
let mut weight: f64 = (0..n).filter(|&v| cur[v]).map(|v| g.nw[v]).sum();
let mut cur_cut = cut(&cur);
let mut best_cut = cur_cut;
let mut best = cur.clone();
let mut gain: Vec<f64> = (0..n)
.map(|v| g.adj[v].iter().map(|&(u, w)| if cur[u] == cur[v] { -w } else { w }).sum())
.collect();
for _ in 0..n {
let mut pick: Option<(f64, usize)> = None;
for v in 0..n {
if locked[v] {
continue;
}
let nwgt = if cur[v] { weight - g.nw[v] } else { weight + g.nw[v] };
if nwgt < lo || nwgt > hi {
continue;
}
if pick.is_none_or(|(gp, _)| gain[v] > gp) {
pick = Some((gain[v], v));
}
}
let Some((gv, v)) = pick else { break };
locked[v] = true;
weight = if cur[v] { weight - g.nw[v] } else { weight + g.nw[v] };
cur[v] = !cur[v];
cur_cut -= gv;
gain[v] = -gv;
for &(u, w) in &g.adj[v] {
gain[u] += if cur[u] == cur[v] { -2.0 * w } else { 2.0 * w };
}
if cur_cut < best_cut - 1e-9 {
best_cut = cur_cut;
best = cur.clone();
}
}
if best == *side {
break;
}
*side = best;
}
}
fn grow(g: &WGraph, target: f64, rng: &mut Rng) -> Vec<bool> {
let n = g.n();
let mut side = vec![false; n];
let mut weight = 0.0;
let mut queue = std::collections::VecDeque::new();
let mut start = rng.below(n);
while weight < target {
if queue.is_empty() {
let mut tries = 0;
while side[start] && tries < n {
start = (start + 1) % n;
tries += 1;
}
if side[start] {
break;
}
side[start] = true;
weight += g.nw[start];
queue.push_back(start);
continue;
}
let v = queue.pop_front().unwrap();
for &(u, _) in &g.adj[v] {
if weight < target && !side[u] {
side[u] = true;
weight += g.nw[u];
queue.push_back(u);
}
}
}
side
}
fn bisect(ids: &[usize], live: &[Option<Vec<u32>>], dims: &[usize], p: &PartitionParams, rng: &mut Rng) -> (Vec<usize>, Vec<usize>) {
let n = ids.len();
let mut by_index: BTreeMap<u32, Vec<usize>> = BTreeMap::new();
for (k, &t) in ids.iter().enumerate() {
for &i in live[t].as_ref().unwrap() {
by_index.entry(i).or_default().push(k);
}
}
let mut adj: Vec<Vec<(usize, f64)>> = vec![Vec::new(); n];
for (&i, ks) in &by_index {
if ks.len() == 2 {
let w = ln(dims[i as usize] as f64) / core::f64::consts::LN_2;
adj[ks[0]].push((ks[1], w));
adj[ks[1]].push((ks[0], w));
}
}
let mut levels = vec![WGraph { nw: vec![1.0; n], adj: merged(adj) }];
let mut maps: Vec<Vec<usize>> = Vec::new();
while levels.last().unwrap().n() > 24 {
let (coarse, map) = coarsen(levels.last().unwrap(), rng);
if coarse.n() as f64 > 0.95 * levels.last().unwrap().n() as f64 {
break;
}
levels.push(coarse);
maps.push(map);
}
let total = n as f64;
let lo = (total * (0.5 - p.imbalance / 2.0)).max(1.0);
let hi = (total * (0.5 + p.imbalance / 2.0)).min(total - 1.0);
let coarsest = levels.last().unwrap();
let mut side: Vec<bool> = Vec::new();
let mut best_cut = f64::INFINITY;
for _ in 0..6 {
let target = lo + (hi - lo) * rng.unit();
let mut s = grow(coarsest, target, rng);
refine(coarsest, &mut s, lo, hi, p.passes);
let w: f64 = (0..coarsest.n()).filter(|&v| s[v]).map(|v| coarsest.nw[v]).sum();
let c: f64 = (0..coarsest.n())
.flat_map(|v| coarsest.adj[v].iter().filter(move |&&(u, _)| v < u).map(move |&(u, w)| (v, u, w)))
.filter(|&(v, u, _)| s[v] != s[u])
.map(|t| t.2)
.sum();
let fits = w >= lo && w <= hi;
if fits && c < best_cut {
best_cut = c;
side = s;
} else if side.is_empty() {
side = s;
}
}
for lvl in (0..maps.len()).rev() {
let map = &maps[lvl];
let mut fine: Vec<bool> = map.iter().map(|&c| side[c]).collect();
refine(&levels[lvl], &mut fine, lo, hi, p.passes);
side = fine;
}
if side.iter().all(|&x| x) || side.iter().all(|&x| !x) {
side = (0..n).map(|k| k < n / 2).collect();
}
let left: Vec<usize> = (0..n).filter(|&k| side[k]).map(|k| ids[k]).collect();
let right: Vec<usize> = (0..n).filter(|&k| !side[k]).map(|k| ids[k]).collect();
(left, right)
}
pub fn reconfigure(shape: &[Vec<u32>], dims: &[usize], open: &[u32], path: &Path, k: usize) -> Path {
let n = shape.len();
let mut sets: Vec<Vec<u32>> = shape.to_vec();
let mut kids: Vec<Option<(usize, usize)>> = vec![None; n];
for &(a, b) in &path.steps {
sets.push(sym_diff(&sets[a], &sets[b], open));
kids.push(Some((a, b)));
}
if kids.len() <= n {
return path.clone();
}
let flops = |a: &[u32], b: &[u32]| set_size(&union(a, b), dims);
let root = kids.len() - 1;
let mut order = vec![root];
let mut i = 0;
while i < order.len() {
if let Some((a, b)) = kids[order[i]] {
order.push(a);
order.push(b);
}
i += 1;
}
for &v in &order {
if kids[v].is_none() {
continue;
}
let mut frontier = vec![v];
loop {
let mut pick: Option<(f64, usize)> = None;
for (fi, &f) in frontier.iter().enumerate() {
if let Some((a, b)) = kids[f] {
let c = flops(&sets[a], &sets[b]);
if pick.is_none_or(|(pc, _)| c > pc) {
pick = Some((c, fi));
}
}
}
let Some((_, fi)) = pick else { break };
if frontier.len() + 1 > k {
break;
}
let f = frontier.swap_remove(fi);
let (a, b) = kids[f].unwrap();
frontier.push(a);
frontier.push(b);
}
let m = frontier.len();
if m < 3 {
continue;
}
let mut current = 0.0;
let mut stack = vec![v];
while let Some(x) = stack.pop() {
if frontier.contains(&x) {
continue;
}
if let Some((a, b)) = kids[x] {
current += flops(&sets[a], &sets[b]);
stack.push(a);
stack.push(b);
}
}
let full = (1usize << m) - 1;
let mut out: Vec<Vec<u32>> = vec![Vec::new(); full + 1];
let mut cost = vec![f64::INFINITY; full + 1];
let mut split = vec![0usize; full + 1];
for j in 0..m {
out[1 << j] = sets[frontier[j]].clone();
cost[1 << j] = 0.0;
}
for sub in 1..=full {
if sub.count_ones() < 2 {
continue;
}
let low = sub.isolate_lowest_one();
let rest = sub ^ low;
out[sub] = sym_diff(&out[low], &out[rest], open);
let mut a = (sub - 1) & sub;
while a > 0 {
if a & low != 0 {
let b = sub ^ a;
let c = cost[a] + cost[b] + flops(&out[a], &out[b]);
if c < cost[sub] {
cost[sub] = c;
split[sub] = a;
}
}
a = (a - 1) & sub;
}
}
if cost[full] >= current * (1.0 - 1e-12) {
continue;
}
fn build(
sub: usize,
root: Option<usize>,
frontier: &[usize],
split: &[usize],
out: &[Vec<u32>],
sets: &mut Vec<Vec<u32>>,
kids: &mut Vec<Option<(usize, usize)>>,
) -> usize {
if sub.count_ones() == 1 {
return frontier[sub.trailing_zeros() as usize];
}
let a = split[sub];
let l = build(a, None, frontier, split, out, sets, kids);
let r = build(sub ^ a, None, frontier, split, out, sets, kids);
let id = match root {
Some(v) => v,
None => {
sets.push(Vec::new());
kids.push(None);
sets.len() - 1
}
};
sets[id] = out[sub].clone();
kids[id] = Some((l, r));
id
}
build(full, Some(v), &frontier, &split, &out, &mut sets, &mut kids);
}
let mut ssa: Vec<Option<usize>> = vec![None; kids.len()];
for (t, slot) in ssa.iter_mut().enumerate().take(n) {
*slot = Some(t);
}
let mut steps = Vec::new();
let mut next = n;
let mut stack = vec![(root, false)];
while let Some((x, expanded)) = stack.pop() {
let Some((a, b)) = kids[x] else { continue };
if expanded {
steps.push((ssa[a].unwrap(), ssa[b].unwrap()));
ssa[x] = Some(next);
next += 1;
} else {
stack.push((x, true));
stack.push((b, false));
stack.push((a, false));
}
}
Path { steps }
}
fn sliced_dims(dims: &[usize], sliced: &[u32]) -> Vec<usize> {
let mut d = dims.to_vec();
for &i in sliced {
d[i as usize] = 1;
}
d
}
pub fn slice(shape: &[Vec<u32>], dims: &[usize], open: &[u32], path: &Path, max_size: f64) -> Vec<u32> {
let mut sets: Vec<Vec<u32>> = shape.to_vec();
let mut flop_sets: Vec<Vec<u32>> = Vec::with_capacity(path.steps.len());
for &(a, b) in &path.steps {
flop_sets.push(union(&sets[a], &sets[b]));
sets.push(sym_diff(&sets[a], &sets[b], open));
}
let nd = dims.len();
let mut in_flops: Vec<Vec<usize>> = vec![Vec::new(); nd];
for (s, f) in flop_sets.iter().enumerate() {
for &i in f {
in_flops[i as usize].push(s);
}
}
let mut in_sets: Vec<Vec<usize>> = vec![Vec::new(); nd];
for (t, x) in sets.iter().enumerate() {
for &i in x {
in_sets[i as usize].push(t);
}
}
let mut flop_size: Vec<f64> = flop_sets.iter().map(|f| set_size(f, dims)).collect();
let mut set_sz: Vec<f64> = sets.iter().map(|x| set_size(x, dims)).collect();
let mut sliced: Vec<u32> = Vec::new();
let mut slices = 1.0;
let mut is_sliced = vec![false; nd];
loop {
let biggest = set_sz.iter().copied().fold(0.0, f64::max);
if biggest <= max_size {
return sliced;
}
let per: f64 = flop_size.iter().sum();
let mut cands: Vec<u32> = Vec::new();
for (t, x) in sets.iter().enumerate() {
if set_sz[t] > max_size {
cands.extend(x.iter().copied().filter(|i| !is_sliced[*i as usize] && !open.contains(i) && dims[*i as usize] > 1));
}
}
cands.sort_unstable();
cands.dedup();
let mut best: Option<(f64, u32)> = None;
for &i in &cands {
let d = dims[i as usize] as f64;
let saved: f64 = in_flops[i as usize].iter().map(|&s| flop_size[s]).sum::<f64>() * (1.0 - 1.0 / d);
let total = (per - saved) * slices * d;
if best.is_none_or(|(b, _)| total < b) {
best = Some((total, i));
}
}
let Some((_, i)) = best else { return sliced };
let d = dims[i as usize] as f64;
for &s in &in_flops[i as usize] {
flop_size[s] /= d;
}
for &t in &in_sets[i as usize] {
set_sz[t] /= d;
}
is_sliced[i as usize] = true;
slices *= d;
sliced.push(i);
}
}
#[derive(Clone, Copy, Debug)]
pub struct PathOptions {
pub trials: usize,
pub seed: u64,
pub max_size: Option<f64>,
pub reconfigure: usize,
pub families: Families,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Families {
Both,
Greedy,
Partition,
}
impl Default for PathOptions {
fn default() -> Self {
PathOptions { trials: 64, seed: 0, max_size: None, reconfigure: 8, families: Families::Both }
}
}
#[derive(Clone, Debug, PartialEq)]
pub struct Plan {
pub path: Path,
pub sliced: Vec<u32>,
pub per_slice: Cost,
pub slices: f64,
}
impl Plan {
pub fn total_flops(&self) -> f64 {
self.per_slice.flops * self.slices
}
}
pub fn find_path(shape: &[Vec<u32>], dims: &[usize], open: &[u32], opts: &PathOptions) -> Plan {
let mut best = search(shape, dims, open, opts);
if opts.max_size.is_some() && best.slices > 1.0 {
let first = best.sliced.clone();
let mut fixings = vec![first.clone()];
if first.len() > 1 {
fixings.push(first[..first.len().div_ceil(2)].to_vec());
}
for (k, pre) in fixings.into_iter().enumerate() {
let fixed = sliced_dims(dims, &pre);
let again = PathOptions { seed: opts.seed ^ (0x5ee_d0f5_11ce + k as u64), trials: opts.trials.div_ceil(2), ..*opts };
let mut second = search(shape, &fixed, open, &again);
let pre_slices: f64 = pre.iter().map(|&i| dims[i as usize] as f64).product();
let mut sliced = pre;
sliced.extend(second.sliced.iter().copied());
second.sliced = sliced;
second.slices *= pre_slices;
if second.total_flops() < best.total_flops() {
best = second;
}
}
}
best
}
fn search(shape: &[Vec<u32>], dims: &[usize], open: &[u32], opts: &PathOptions) -> Plan {
let mut rng = Rng::new(opts.seed);
let score = |path: &Path| -> Plan {
let sliced = match opts.max_size {
Some(m) => slice(shape, dims, open, path, m),
None => Vec::new(),
};
let d = sliced_dims(dims, &sliced);
let per_slice = path_cost(shape, &d, open, path);
let slices = sliced.iter().map(|&i| dims[i as usize] as f64).product();
Plan { path: path.clone(), sliced, per_slice, slices }
};
const LN_1000: f64 = 6.907_755_278_982_137;
let trials = opts.trials.max(1);
let mut best_greedy: Option<(f64, f64, f64)> = None;
let mut best_part: Option<(f64, PartitionParams)> = None;
let mut top: Vec<Plan> = Vec::new();
let keep = 4;
for t in 0..trials {
let seed = rng.next();
let greedy = match opts.families {
Families::Greedy => true,
Families::Partition => false,
Families::Both => {
if t < 8 {
t % 2 == 0
} else {
let greedy_ahead = match (best_greedy, best_part) {
(Some(g), Some(p)) => g.0 <= p.0,
(Some(_), None) => true,
_ => false,
};
(rng.unit() < 0.75) == greedy_ahead
}
}
};
let local = t >= trials / 2 && rng.unit() < 0.5;
let path;
if greedy {
let (mut costmod, mut temperature) =
(0.1 + 3.9 * rng.unit(), if t == 0 { 0.0 } else { crate::repro::exp(-LN_1000 * rng.unit()) });
if let (true, Some((_, c, tp))) = (local, best_greedy) {
costmod = (c + 0.6 * (rng.unit() - 0.5)).clamp(0.1, 4.0);
temperature = (tp.max(1e-3) * crate::repro::exp(2.0 * rng.unit() - 1.0)).min(1.0);
}
path = greedy_path(shape, dims, open, costmod, temperature, seed);
let plan = score(&path);
if best_greedy.is_none_or(|b| plan.total_flops() < b.0) {
best_greedy = Some((plan.total_flops(), costmod, temperature));
}
push_top(&mut top, plan, keep);
} else {
let mut p = PartitionParams {
imbalance: 0.02 + 0.6 * rng.unit(),
cutoff: 4 + rng.below(12),
passes: 2 + rng.below(4),
costmod: 0.1 + 3.9 * rng.unit(),
temperature: crate::repro::exp(-LN_1000 * rng.unit()),
};
if let (true, Some((_, b))) = (local, best_part) {
p.imbalance = (b.imbalance + 0.2 * (rng.unit() - 0.5)).clamp(0.01, 0.9);
p.cutoff = (b.cutoff as isize + rng.below(5) as isize - 2).clamp(3, 20) as usize;
p.costmod = (b.costmod + 0.6 * (rng.unit() - 0.5)).clamp(0.1, 4.0);
p.temperature = b.temperature;
}
path = partition_path(shape, dims, open, p, seed);
let plan = score(&path);
if best_part.is_none_or(|b| plan.total_flops() < b.0) {
best_part = Some((plan.total_flops(), p));
}
push_top(&mut top, plan, keep);
}
}
let mut best = top[0].clone();
if opts.reconfigure >= 3 {
for cand in top {
let mut cur = cand;
for _ in 0..8 {
let d = sliced_dims(dims, &cur.sliced);
let plan = score(&reconfigure(shape, &d, open, &cur.path, opts.reconfigure));
if plan.total_flops() < cur.total_flops() * (1.0 - 1e-9) {
cur = plan;
} else {
break;
}
}
if cur.total_flops() < best.total_flops() {
best = cur;
}
}
}
best
}
fn push_top(top: &mut Vec<Plan>, plan: Plan, keep: usize) {
let at = top.iter().position(|p| plan.total_flops() < p.total_flops()).unwrap_or(top.len());
if at < keep {
top.insert(at, plan);
top.truncate(keep);
}
}
pub fn contract_path(net: &Network, path: &Path) -> Result<Tensor, TnError> {
let n = net.tensors.len();
if path.steps.len() + 1 != n.max(1) {
return Err(TnError::PlanMismatch);
}
let mut live: Vec<Option<Tensor>> = net.tensors.iter().cloned().map(Some).collect();
for &(a, b) in &path.steps {
let ta = live.get_mut(a).and_then(Option::take).ok_or(TnError::PlanMismatch)?;
let tb = live.get_mut(b).and_then(Option::take).ok_or(TnError::PlanMismatch)?;
live.push(Some(contract(&ta, &tb, &net.dims, &net.open)));
}
let last = live.into_iter().rev().flatten().next().ok_or(TnError::PlanMismatch)?;
let order = net.open.clone();
let data = permute(&last, &order, &net.dims);
Ok(Tensor { inds: order, data })
}
fn fix(net: &Network, sliced: &[u32], values: &[usize]) -> Network {
let mut tensors = Vec::with_capacity(net.tensors.len());
for t in &net.tensors {
let mut cur = t.clone();
for (si, &i) in sliced.iter().enumerate() {
if let Some(pos) = cur.inds.iter().position(|&j| j == i) {
let r = cur.inds.len();
let dim = net.dims[i as usize];
let inner: usize = cur.inds[pos + 1..].iter().map(|&j| net.dims[j as usize]).product();
let outer: usize = cur.inds[..pos].iter().map(|&j| net.dims[j as usize]).product();
let mut data = Vec::with_capacity(outer * inner);
for o in 0..outer {
let base = (o * dim + values[si]) * inner;
data.extend_from_slice(&cur.data[base..base + inner]);
}
let mut inds = cur.inds.clone();
inds.remove(pos);
let _ = r;
cur = Tensor { inds, data };
}
}
tensors.push(cur);
}
Network { tensors, dims: net.dims.clone(), open: net.open.clone() }
}
#[cfg(not(target_arch = "wasm32"))]
const SLICE_MEMORY: f64 = 8.0 * 1024.0 * 1024.0 * 1024.0;
pub fn contract_plan(net: &Network, plan: &Plan) -> Result<Tensor, TnError> {
let dims: Vec<usize> = plan.sliced.iter().map(|&i| net.dims[i as usize]).collect();
let count: usize = dims.iter().product();
let assignment = |mut s: usize| -> Vec<usize> {
let mut v = vec![0; dims.len()];
for k in (0..dims.len()).rev() {
v[k] = s % dims[k];
s /= dims[k];
}
v
};
let run = |s: usize| contract_path(&fix(net, &plan.sliced, &assignment(s)), &plan.path);
#[cfg(not(target_arch = "wasm32"))]
let parts: Vec<Result<Tensor, TnError>> = {
let per_thread = 6.0 * plan.per_slice.max_size * core::mem::size_of::<C>() as f64;
let fit = (SLICE_MEMORY / per_thread).floor().max(1.0) as usize;
let threads = std::thread::available_parallelism().map_or(1, |n| n.get()).min(count).min(fit).max(1);
let mut parts: Vec<Option<Result<Tensor, TnError>>> = (0..count).map(|_| None).collect();
std::thread::scope(|sc| {
let chunks: Vec<_> = parts.chunks_mut(count.div_ceil(threads)).enumerate().collect();
for (c, chunk) in chunks {
let run = &run;
let base = c * count.div_ceil(threads);
sc.spawn(move || {
for (k, slot) in chunk.iter_mut().enumerate() {
*slot = Some(run(base + k));
}
});
}
});
parts.into_iter().map(|p| p.expect("every slice ran")).collect()
};
#[cfg(target_arch = "wasm32")]
let parts: Vec<Result<Tensor, TnError>> = (0..count).map(run).collect();
let mut total: Option<Tensor> = None;
for p in parts {
let t = p?;
total = Some(match total {
None => t,
Some(mut acc) => {
for (x, y) in acc.data.iter_mut().zip(&t.data) {
x.re += y.re;
x.im += y.im;
}
acc
}
});
}
total.ok_or(TnError::PlanMismatch)
}
#[cfg(test)]
pub(crate) mod tests {
use super::*;
#[allow(clippy::needless_range_loop)]
pub(crate) fn random_unitary(k: usize, rng: &mut Rng) -> Vec<C> {
let d = 1 << k;
let mut cols: Vec<Vec<C>> = (0..d)
.map(|_| (0..d).map(|_| C::new(rng.unit() - 0.5, rng.unit() - 0.5)).collect())
.collect();
for j in 0..d {
for i in 0..j {
let mut dot = C::ZERO;
for r in 0..d {
dot = dot.add(cols[i][r].conj().mul(cols[j][r]));
}
for r in 0..d {
let t = cols[i][r].mul(dot);
cols[j][r] = cols[j][r].sub(t);
}
}
let norm = cols[j].iter().map(|c| c.norm2()).sum::<f64>().sqrt();
for r in 0..d {
cols[j][r] = cols[j][r].scale(1.0 / norm);
}
}
let mut m = vec![C::ZERO; d * d];
for r in 0..d {
for c in 0..d {
m[r * d + c] = cols[c][r];
}
}
m
}
pub(crate) fn dense(n: u32, gates: &[Gate]) -> Vec<C> {
let dim = 1usize << n;
let mut psi = vec![C::ZERO; dim];
psi[0] = C::ONE;
for g in gates {
let k = g.qubits.len();
let mut out = vec![C::ZERO; dim];
for (x, amp) in psi.iter().enumerate() {
if amp.re == 0.0 && amp.im == 0.0 {
continue;
}
let col = g.qubits.iter().fold(0usize, |c, &q| (c << 1) | ((x >> (n - 1 - q)) & 1));
for row in 0..1usize << k {
let mut y = x;
for (i, &q) in g.qubits.iter().enumerate() {
let bit = (row >> (k - 1 - i)) & 1;
let shift = n - 1 - q;
y = (y & !(1 << shift)) | (bit << shift);
}
out[y] = out[y].add(g.matrix[row * (1 << k) + col].mul(*amp));
}
}
psi = out;
}
psi
}
pub(crate) fn random_circuit(n: u32, depth: usize, seed: u64) -> Vec<Gate> {
let mut rng = Rng::new(seed);
let mut gates = Vec::new();
for layer in 0..depth {
for q in 0..n {
gates.push(Gate { qubits: vec![q], matrix: random_unitary(1, &mut rng) });
}
let mut q = (layer % 2) as u32;
while q + 1 < n {
gates.push(Gate { qubits: vec![q, q + 1], matrix: random_unitary(2, &mut rng) });
q += 2;
}
let a = rng.below(n as usize) as u32;
let b = rng.below(n as usize) as u32;
if a != b {
gates.push(Gate { qubits: vec![b, a], matrix: random_unitary(2, &mut rng) });
}
}
gates
}
fn close(a: C, b: C) -> bool {
(a.re - b.re).abs() < 1e-12 && (a.im - b.im).abs() < 1e-12
}
#[test]
fn contraction_matches_the_state_vector() {
for seed in 0..6 {
let n = 7;
let gates = random_circuit(n, 6, seed);
let psi = dense(n, &gates);
for x in [0usize, 5, 77, 127] {
let bits: Vec<u8> = (0..n).map(|q| ((x >> (n - 1 - q)) & 1) as u8).collect();
let mut net = Network::amplitude(n, &gates, &bits, &[]).unwrap();
net.simplify();
let path = greedy_path(&net.shape(), &net.dims, &net.open, 1.0, 0.0, 0);
let t = contract_path(&net, &path).unwrap();
assert_eq!(t.data.len(), 1);
assert!(close(t.data[0], psi[x]), "seed {seed} x {x}: {:?} vs {:?}", t.data[0], psi[x]);
}
}
}
#[test]
fn open_outputs_give_a_batch_of_amplitudes() {
let n = 6;
let gates = random_circuit(n, 5, 9);
let psi = dense(n, &gates);
let bits = [1u8, 0, 0, 0, 1, 0];
let mut net = Network::amplitude(n, &gates, &bits, &[5, 1]).unwrap();
net.simplify();
let path = greedy_path(&net.shape(), &net.dims, &net.open, 1.0, 0.3, 4);
let t = contract_path(&net, &path).unwrap();
assert_eq!(t.data.len(), 4);
for (j, got) in t.data.iter().enumerate() {
let (b5, b1) = ((j >> 1) & 1, j & 1);
let x = (1 << 5) | (b1 << 4) | (1 << 1) | b5;
assert!(close(*got, psi[x]), "j {j}");
}
}
#[test]
fn every_path_gives_the_same_amplitude_and_the_same_bits_twice() {
let n = 8;
let gates = random_circuit(n, 6, 3);
let bits = [0u8, 1, 1, 0, 1, 0, 0, 1];
let mut net = Network::amplitude(n, &gates, &bits, &[]).unwrap();
net.simplify();
let reference = contract_path(&net, &greedy_path(&net.shape(), &net.dims, &net.open, 1.0, 0.0, 0)).unwrap().data[0];
for seed in 0..8 {
let p = greedy_path(&net.shape(), &net.dims, &net.open, 0.5 + 0.2 * seed as f64, 0.5, seed);
let a = contract_path(&net, &p).unwrap().data[0];
assert!(close(a, reference));
let b = contract_path(&net, &p).unwrap().data[0];
assert_eq!((a.re.to_bits(), a.im.to_bits()), (b.re.to_bits(), b.im.to_bits()));
}
}
pub(crate) fn grid_circuit(w: u32, h: u32, depth: usize, seed: u64) -> Vec<Gate> {
let mut rng = Rng::new(seed);
let q = |x: u32, y: u32| y * w + x;
let mut gates = Vec::new();
for layer in 0..depth {
for i in 0..w * h {
gates.push(Gate { qubits: vec![i], matrix: random_unitary(1, &mut rng) });
}
let (horizontal, offset) = ((layer / 2) % 2 == 0, (layer % 2) as u32);
if horizontal {
for y in 0..h {
let mut x = offset;
while x + 1 < w {
gates.push(Gate { qubits: vec![q(x, y), q(x + 1, y)], matrix: random_unitary(2, &mut rng) });
x += 2;
}
}
} else {
for x in 0..w {
let mut y = offset;
while y + 1 < h {
gates.push(Gate { qubits: vec![q(x, y), q(x, y + 1)], matrix: random_unitary(2, &mut rng) });
y += 2;
}
}
}
}
gates
}
#[test]
fn every_tree_family_gives_the_amplitude() {
let (w, h) = (3, 3);
let n = w * h;
let gates = grid_circuit(w, h, 8, 21);
let psi = dense(n, &gates);
let bits: Vec<u8> = (0..n).map(|q| (q % 3 == 0) as u8).collect();
let x = bits.iter().fold(0usize, |a, &b| (a << 1) | b as usize);
let mut net = Network::amplitude(n, &gates, &bits, &[]).unwrap();
net.simplify();
let shape = net.shape();
for seed in 0..6 {
let p = PartitionParams { imbalance: 0.3, cutoff: 3 + seed as usize, passes: 3, costmod: 1.0, temperature: 0.1 };
let path = partition_path(&shape, &net.dims, &net.open, p, seed);
assert!(close(contract_path(&net, &path).unwrap().data[0], psi[x]), "partition seed {seed}");
let r = reconfigure(&shape, &net.dims, &net.open, &path, 8);
let (c0, c1) = (path_cost(&shape, &net.dims, &net.open, &path), path_cost(&shape, &net.dims, &net.open, &r));
assert!(c1.flops <= c0.flops, "reconfiguration made it costlier");
assert!(close(contract_path(&net, &r).unwrap().data[0], psi[x]), "reconfigured seed {seed}");
}
}
#[test]
fn hyper_optimisation_beats_plain_greedy() {
let gates = grid_circuit(4, 4, 10, 2);
let n = 16;
let mut net = Network::amplitude(n, &gates, &vec![0; n as usize], &[]).unwrap();
net.simplify();
let shape = net.shape();
let plain = path_cost(&shape, &net.dims, &net.open, &greedy_path(&shape, &net.dims, &net.open, 1.0, 0.0, 0));
let plan = find_path(&shape, &net.dims, &net.open, &PathOptions { trials: 24, seed: 1, max_size: None, reconfigure: 6, families: Families::Both });
assert!(plan.total_flops() < plain.flops, "{} vs {}", plan.total_flops(), plain.flops);
let again = find_path(&shape, &net.dims, &net.open, &PathOptions { trials: 24, seed: 1, max_size: None, reconfigure: 6, families: Families::Both });
assert_eq!(plan, again);
}
#[test]
fn slices_fit_the_cap_and_sum_to_the_amplitude() {
let (w, h) = (4, 3);
let n = w * h;
let gates = grid_circuit(w, h, 10, 5);
let psi = dense(n, &gates);
let bits: Vec<u8> = (0..n).map(|q| (q % 2) as u8).collect();
let x = bits.iter().fold(0usize, |a, &b| (a << 1) | b as usize);
let mut net = Network::amplitude(n, &gates, &bits, &[]).unwrap();
net.simplify();
let shape = net.shape();
let free = find_path(&shape, &net.dims, &net.open, &PathOptions { trials: 12, seed: 3, max_size: None, reconfigure: 6, families: Families::Both });
let cap = free.per_slice.max_size / 16.0;
let plan = find_path(&shape, &net.dims, &net.open, &PathOptions { trials: 12, seed: 3, max_size: Some(cap), reconfigure: 6, families: Families::Both });
assert!(plan.per_slice.max_size <= cap);
assert!(plan.slices >= 2.0);
let t = contract_plan(&net, &plan).unwrap();
assert!(close(t.data[0], psi[x]), "{:?} vs {:?}", t.data[0], psi[x]);
let u = contract_plan(&net, &plan).unwrap();
assert_eq!((t.data[0].re.to_bits(), t.data[0].im.to_bits()), (u.data[0].re.to_bits(), u.data[0].im.to_bits()));
}
#[test]
fn reconfiguration_reaches_the_optimum_on_small_networks() {
fn optimum(shape: &[Vec<u32>], dims: &[usize], set: u32, memo: &mut BTreeMap<u32, (f64, Vec<u32>)>) -> (f64, Vec<u32>) {
if set.count_ones() == 1 {
return (0.0, shape[set.trailing_zeros() as usize].clone());
}
if let Some(v) = memo.get(&set) {
return v.clone();
}
let mut best: Option<(f64, Vec<u32>)> = None;
let mut a = (set - 1) & set;
while a > 0 {
let b = set ^ a;
if a < b {
let (ca, sa) = optimum(shape, dims, a, memo);
let (cb, sb) = optimum(shape, dims, b, memo);
let c = ca + cb + set_size(&union(&sa, &sb), dims);
if best.as_ref().is_none_or(|x| c < x.0) {
best = Some((c, sym_diff(&sa, &sb, &[])));
}
}
a = (a - 1) & set;
}
let v = best.unwrap();
memo.insert(set, v.clone());
v
}
let mut rng = Rng::new(77);
for trial in 0..40 {
let n = 5 + rng.below(4); let mut shape: Vec<Vec<u32>> = vec![Vec::new(); n];
let mut dims = Vec::new();
let bond = |a: usize, b: usize, shape: &mut Vec<Vec<u32>>, dims: &mut Vec<usize>, rng: &mut Rng| {
let i = dims.len() as u32;
dims.push(2 + rng.below(3));
shape[a].push(i);
shape[b].push(i);
};
for t in 1..n {
let u = rng.below(t);
bond(u, t, &mut shape, &mut dims, &mut rng);
}
for _ in 0..n {
let (a, b) = (rng.below(n), rng.below(n));
if a != b {
bond(a, b, &mut shape, &mut dims, &mut rng);
}
}
let path = greedy_path(&shape, &dims, &[], 1.0, 1.0, trial);
let r = reconfigure(&shape, &dims, &[], &path, 8);
let got = path_cost(&shape, &dims, &[], &r).flops;
let (want, _) = optimum(&shape, &dims, (1u32 << n) - 1, &mut BTreeMap::new());
assert!((got - want).abs() <= 1e-9 * want, "trial {trial}: reconfigured {got}, optimum {want}");
}
}
#[test]
fn slicing_prices_match_a_full_recount() {
for (seed, cap_div) in [(1u64, 8.0), (2, 64.0), (3, 512.0)] {
let gates = grid_circuit(4, 4, 9, seed);
let mut net = Network::amplitude(16, &gates, &[0; 16], &[]).unwrap();
net.simplify();
let (shape, dims, open) = (net.shape(), net.dims.clone(), net.open.clone());
let path = greedy_path(&shape, &dims, &open, 2.0, 0.2, seed);
let cap = path_cost(&shape, &dims, &open, &path).max_size / cap_div;
let mut want: Vec<u32> = Vec::new();
loop {
let d = sliced_dims(&dims, &want);
if path_cost(&shape, &d, &open, &path).max_size <= cap {
break;
}
let mut sets: Vec<Vec<u32>> = shape.clone();
for &(a, b) in &path.steps {
sets.push(sym_diff(&sets[a], &sets[b], &open));
}
let mut cands: Vec<u32> =
sets.iter().filter(|x| set_size(x, &d) > cap).flatten().copied().filter(|i| d[*i as usize] > 1).collect();
cands.sort_unstable();
cands.dedup();
let mut best: Option<(f64, u32)> = None;
for &i in &cands {
let mut t = want.clone();
t.push(i);
let total = path_cost(&shape, &sliced_dims(&dims, &t), &open, &path).flops * (1u64 << t.len()) as f64;
if best.is_none_or(|(b, _)| total < b) {
best = Some((total, i));
}
}
want.push(best.unwrap().1);
}
assert_eq!(slice(&shape, &dims, &open, &path, cap), want, "seed {seed}");
}
}
#[test]
fn bad_input_is_refused() {
let g = vec![Gate { qubits: vec![0, 0], matrix: vec![C::ONE; 16] }];
assert_eq!(Network::amplitude(2, &g, &[0, 0], &[]).unwrap_err(), TnError::BadGate(0));
assert_eq!(Network::amplitude(2, &[], &[0], &[]).unwrap_err(), TnError::BadOutput);
assert_eq!(Network::amplitude(2, &[], &[0, 2], &[]).unwrap_err(), TnError::BadOutput);
}
}