use super::simplify::OriginalFate;
use crate::cnf::{Original, Reduced, ShowSet, Space, VarId, Weights};
use serde::de::{self, Deserializer, Visitor};
use serde::{Deserialize, Serialize, Serializer};
use std::marker::PhantomData;
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(transparent, bound = "")]
pub struct VarMap<Src: Space, Tgt: Space> {
entries: Vec<Option<i32>>,
#[serde(skip)]
spaces: PhantomData<(Src, Tgt)>,
}
impl<Src: Space, Tgt: Space> VarMap<Src, Tgt> {
pub fn from_entries(entries: Vec<Option<i32>>) -> Self {
VarMap {
entries,
spaces: PhantomData,
}
}
pub fn identity(num_vars: u32) -> Self {
VarMap::from_entries((1..=num_vars as i32).map(Some).collect())
}
pub fn len(&self) -> usize {
self.entries.len()
}
pub fn is_empty(&self) -> bool {
self.entries.is_empty()
}
pub fn get(&self, source_var: VarId) -> Option<i32> {
self.entries.get(source_var.idx()).copied().flatten()
}
pub fn iter(&self) -> impl ExactSizeIterator<Item = Option<i32>> + '_ {
self.entries.iter().copied()
}
pub fn is_injective(&self, target_num_vars: u32) -> bool {
let nv = target_num_vars as usize;
let mut claimed = vec![false; nv];
for e in self.entries.iter().flatten() {
let Some(t) = VarId::try_from_dimacs(*e).map(VarId::idx) else {
return false;
};
if t >= nv || claimed[t] {
return false;
}
claimed[t] = true;
}
true
}
pub fn invert(&self, target_num_vars: u32) -> VarMap<Tgt, Src> {
self.invert_composed(target_num_vars, |source_var| {
VarId(source_var as u32).to_dimacs()
})
}
pub fn carry_show(&self, target_show: &ShowSet<Tgt>) -> ShowSet<Src> {
ShowSet::from_zero_based(
self.entries
.iter()
.enumerate()
.filter_map(|(source, entry)| {
let target = VarId::try_from_dimacs(*entry.as_ref()?)?;
target_show.contains(target).then_some(source as u32)
}),
)
}
pub fn carry_weights(&self, target: &Weights<Tgt>) -> Weights<Src> {
Weights::from_carried(self.entries.iter().map(|entry| {
let lit = (*entry)?;
let (wn, wp) = target.get(VarId::try_from_dimacs(lit)?)?;
Some(if lit > 0 {
(wn.clone(), wp.clone())
} else {
(wp.clone(), wn.clone())
})
}))
}
pub(crate) fn invert_composed<Named: Space>(
&self,
target_num_vars: u32,
source_dimacs: impl Fn(usize) -> i32,
) -> VarMap<Tgt, Named> {
let nv = target_num_vars as usize;
let mut out: Vec<Option<i32>> = vec![None; nv];
for (source_var, entry) in self.entries.iter().enumerate() {
let Some(lit) = entry else { continue };
let Some(target) = VarId::try_from_dimacs(*lit).map(VarId::idx) else {
continue;
};
debug_assert!(
target < nv && out[target].is_none(),
"the map must be injective into the target space before it is inverted",
);
if target >= nv {
continue;
}
let named = source_dimacs(source_var);
out[target] = Some(if *lit > 0 { named } else { -named });
}
VarMap::from_entries(out)
}
}
impl<Src: Space> VarMap<Src, Reduced> {
pub(crate) fn assume_original_target(self) -> VarMap<Src, Original> {
VarMap::from_entries(self.entries)
}
}
impl<Src: Space, Tgt: Space> FromIterator<Option<i32>> for VarMap<Src, Tgt> {
fn from_iter<T: IntoIterator<Item = Option<i32>>>(iter: T) -> Self {
VarMap::from_entries(iter.into_iter().collect())
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OriginalTarget {
Literal(i32),
Constant(bool),
Free,
}
impl Serialize for OriginalTarget {
fn serialize<S: Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
match *self {
OriginalTarget::Literal(lit) => s.serialize_i32(lit),
OriginalTarget::Constant(value) => s.serialize_bool(value),
OriginalTarget::Free => s.serialize_none(),
}
}
}
impl<'de> Deserialize<'de> for OriginalTarget {
fn deserialize<D: Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
struct TargetVisitor;
impl<'de> Visitor<'de> for TargetVisitor {
type Value = OriginalTarget;
fn expecting(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
f.write_str("a signed reduced literal, a boolean, or null")
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<Self::Value, E> {
match i32::try_from(v) {
Ok(lit) => Ok(OriginalTarget::Literal(lit)),
Err(_) => Err(E::invalid_value(de::Unexpected::Signed(v), &self)),
}
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<Self::Value, E> {
match i32::try_from(v) {
Ok(lit) => Ok(OriginalTarget::Literal(lit)),
Err(_) => Err(E::invalid_value(de::Unexpected::Unsigned(v), &self)),
}
}
fn visit_bool<E: de::Error>(self, v: bool) -> Result<Self::Value, E> {
Ok(OriginalTarget::Constant(v))
}
fn visit_unit<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(OriginalTarget::Free)
}
fn visit_none<E: de::Error>(self) -> Result<Self::Value, E> {
Ok(OriginalTarget::Free)
}
}
d.deserialize_any(TargetVisitor)
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct OriginalMap(Vec<OriginalTarget>);
impl OriginalMap {
pub fn from_entries(entries: Vec<OriginalTarget>) -> Self {
OriginalMap(entries)
}
pub fn identity(num_vars: u32) -> Self {
OriginalMap((1..=num_vars as i32).map(OriginalTarget::Literal).collect())
}
pub(crate) fn from_fates(fates: &[OriginalFate]) -> Self {
OriginalMap(
fates
.iter()
.map(|fate| match *fate {
OriginalFate::Variable {
index,
same_polarity,
} => {
let dimacs = VarId(index as u32).to_dimacs();
OriginalTarget::Literal(if same_polarity { dimacs } else { -dimacs })
}
OriginalFate::Forced(value) => OriginalTarget::Constant(value),
OriginalFate::Unconstrained => OriginalTarget::Free,
})
.collect(),
)
}
pub fn len(&self) -> usize {
self.0.len()
}
pub fn is_empty(&self) -> bool {
self.0.is_empty()
}
pub fn get(&self, original_var: VarId) -> Option<OriginalTarget> {
self.0.get(original_var.idx()).copied()
}
pub fn iter(&self) -> impl ExactSizeIterator<Item = OriginalTarget> + '_ {
self.0.iter().copied()
}
}