#include "arjun_shim.h"
#include "arjun.h"
#include <memory>
#include <sstream>
#include <stdexcept>
#include <string>
#include <vector>
#include <cerrno>
#include <climits>
#include <cstdint>
#include <cstring>
#include <cstdlib>
#include <cstdio>
using ArjunNS::Arjun;
using ArjunNS::SimplifiedCNF;
using ArjunNS::SimpConf;
using ArjunNS::FGenMpz;
using ArjunNS::FGenMpq;
using CMSat::Field;
using CMSat::FieldGen;
using CMSat::Lit;
static bool parse_whole_number(const char *v, long lo, long hi, long *out) {
if (v == nullptr || *v == '\0') return false;
char *end = nullptr;
errno = 0;
const long parsed = strtol(v, &end, 10);
if (errno != 0 || end == v || *end != '\0') return false;
if (parsed < lo || parsed > hi) return false;
*out = parsed;
return true;
}
struct ArjunShim {
std::unique_ptr<FieldGen> fg;
std::unique_ptr<Arjun> arjun;
std::unique_ptr<SimplifiedCNF> cur; uint32_t verb = 0;
int64_t backbone_max_confl = -1;
double oracle_mult = -1.0;
std::vector<int32_t> backbone_dimacs;
std::vector<int32_t> eq_dimacs;
ArjunShim(bool weighted, uint32_t seed)
: fg(weighted ? std::unique_ptr<FieldGen>(new FGenMpq)
: std::unique_ptr<FieldGen>(new FGenMpz)),
arjun(new Arjun),
cur(std::make_unique<SimplifiedCNF>(fg)) {
arjun->set_verb(0);
arjun->set_seed(seed);
if (weighted) cur->set_weighted(true);
}
};
static inline int32_t lit_to_dimacs(const Lit &l) {
int32_t d = (int32_t)(l.var() + 1u);
return l.sign() ? -d : d;
}
extern "C" {
ArjunShim *arjun_shim_new(uint32_t seed) {
try {
return new ArjunShim(false, seed);
} catch (...) {
return nullptr;
}
}
ArjunShim *arjun_shim_new_weighted(uint32_t seed) {
try {
return new ArjunShim(true, seed);
} catch (...) {
return nullptr;
}
}
void arjun_shim_free(ArjunShim *s) { delete s; }
void arjun_shim_new_vars(ArjunShim *s, uint32_t n) { s->cur->new_vars(n); }
void arjun_shim_add_clause(ArjunShim *s, const int32_t *lits, size_t n) {
std::vector<Lit> cl;
cl.reserve(n);
for (size_t i = 0; i < n; i++) {
int32_t l = lits[i];
uint32_t var = (uint32_t)(l < 0 ? -l : l) - 1u; cl.push_back(Lit(var, l < 0));
}
s->cur->add_clause(cl);
}
void arjun_shim_set_sampl(ArjunShim *s, const uint32_t *vars0, size_t n) {
std::vector<uint32_t> v(vars0, vars0 + n);
s->cur->set_sampl_vars(v);
}
void arjun_shim_set_verb(ArjunShim *s, uint32_t verb) {
s->verb = verb;
s->arjun->set_verb(verb);
}
void arjun_shim_set_backbone_max_confl(ArjunShim *s, int64_t max_confl) {
s->backbone_max_confl = max_confl;
}
void arjun_shim_set_oracle_mult(ArjunShim *s, double mult) {
s->oracle_mult = mult;
}
void arjun_shim_set_deadline_ms(ArjunShim *s, int64_t ms_from_now) {
s->arjun->set_deadline(ms_from_now < 0 ? -1.0 : (double)ms_from_now / 1000.0);
}
int arjun_shim_stage_minimize_indep(ArjunShim *s, int all_indep) {
try {
ArjunNS::Arjun::IndepInfo info =
s->arjun->standalone_minimize_indep_info(*s->cur, all_indep != 0);
s->backbone_dimacs.clear();
s->backbone_dimacs.reserve(info.backbone.size());
for (const auto &l : info.backbone) s->backbone_dimacs.push_back(lit_to_dimacs(l));
s->eq_dimacs.clear();
s->eq_dimacs.reserve(info.eq_lits.size() * 2);
for (const auto &p : info.eq_lits) {
s->eq_dimacs.push_back(lit_to_dimacs(p.first));
s->eq_dimacs.push_back(lit_to_dimacs(p.second));
}
return 0;
} catch (const std::exception& e) {
if (s->verb > 0) fprintf(stderr, "[arjun-shim] stage_minimize_indep: %s\n", e.what());
return 1;
} catch (...) {
if (s->verb > 0) fprintf(stderr, "[arjun-shim] stage_minimize_indep: unknown failure\n");
return 1;
}
}
int arjun_shim_stage_simplify(ArjunShim *s, int all_indep, int oracle_enabled, int no_sbva, int no_bve) {
try {
Arjun::ElimToFileConf etof;
etof.all_indep = all_indep != 0;
if (no_sbva) {
etof.num_sbva_steps = 0;
}
SimpConf sc;
sc.backbone_max_confl = s->backbone_max_confl;
if (s->oracle_mult >= 0.0) sc.oracle_mult = s->oracle_mult;
if (no_bve || getenv("VITRI_ARJUN_NO_BVE")) sc.do_bve = false;
if (const char *e = getenv("VITRI_ARJUN_BVE_GROW")) {
long g = 0;
if (!parse_whole_number(e, 0, (long)INT_MAX, &g)) {
if (s->verb > 0)
fprintf(stderr,
"[arjun-shim] VITRI_ARJUN_BVE_GROW is not a clause-growth budget\n");
return 1;
}
sc.bve_grow_iter1 = (int)g;
sc.bve_grow_iter2 = (int)g;
}
if (oracle_enabled == 0 || getenv("VITRI_ARJUN_NO_ORACLE")) {
sc.oracle_extra = false;
sc.oracle_vivify = false;
sc.oracle_vivify_get_learnts = false;
sc.oracle_sparsify = false;
}
s->arjun->standalone_elim_to_file(*s->cur, etof, sc);
return 0;
} catch (const std::exception& e) {
if (s->verb > 0) fprintf(stderr, "[arjun-shim] stage_simplify: %s\n", e.what());
return 1;
} catch (...) {
if (s->verb > 0) fprintf(stderr, "[arjun-shim] stage_simplify: unknown failure\n");
return 1;
}
}
uint32_t arjun_shim_cur_nvars(ArjunShim *s) { return s->cur->nVars(); }
size_t arjun_shim_cur_nclauses(ArjunShim *s) { return s->cur->get_clauses().size(); }
size_t arjun_shim_cur_clauses(ArjunShim *s, int32_t *buf, size_t cap) {
const auto &cls = s->cur->get_clauses();
size_t need = 0;
for (const auto &cl : cls) need += cl.size() + 1; if (cap < need || buf == nullptr) return need;
size_t k = 0;
for (const auto &cl : cls) {
for (const auto &l : cl) {
int32_t lit = (int32_t)(l.var() + 1);
buf[k++] = l.sign() ? -lit : lit;
}
buf[k++] = 0;
}
return need;
}
size_t arjun_shim_cur_sampl(ArjunShim *s, uint32_t *buf, size_t cap) {
const auto &sv = s->cur->get_sampl_vars();
if (cap < sv.size() || buf == nullptr) return sv.size();
std::memcpy(buf, sv.data(), sv.size() * sizeof(uint32_t));
return sv.size();
}
size_t arjun_shim_backbone(ArjunShim *s, int32_t *buf, size_t cap) {
const auto &bb = s->backbone_dimacs;
if (cap < bb.size() || buf == nullptr) return bb.size();
std::memcpy(buf, bb.data(), bb.size() * sizeof(int32_t));
return bb.size();
}
size_t arjun_shim_eq_lits(ArjunShim *s, int32_t *buf, size_t cap) {
const auto &eq = s->eq_dimacs; if (cap < eq.size() || buf == nullptr) return eq.size();
std::memcpy(buf, eq.data(), eq.size() * sizeof(int32_t));
return eq.size();
}
static const size_t RED_CLAUSE_MAX_LEN = 8; static const size_t RED_CLAUSE_MAX_TOTAL = 50000;
size_t arjun_shim_red_clauses(ArjunShim *s, int32_t *buf, size_t cap) {
const auto &cls = s->cur->get_red_clauses();
size_t need = 0, kept = 0;
for (const auto &cl : cls) {
if (kept >= RED_CLAUSE_MAX_TOTAL) break;
if (cl.size() > RED_CLAUSE_MAX_LEN) continue; need += cl.size() + 1;
kept++;
}
if (cap < need || buf == nullptr) return need;
size_t k = 0, kept2 = 0;
for (const auto &cl : cls) {
if (kept2 >= RED_CLAUSE_MAX_TOTAL) break;
if (cl.size() > RED_CLAUSE_MAX_LEN) continue;
for (const auto &l : cl) {
int32_t lit = (int32_t)(l.var() + 1); buf[k++] = l.sign() ? -lit : lit;
}
buf[k++] = 0;
kept2++;
}
return need;
}
size_t arjun_shim_orig_to_new(ArjunShim *s, int32_t *buf, size_t cap) {
const auto &m = s->cur->get_orig_to_new_var();
size_t need = 0;
for (const auto &kv : m) {
if (kv.second == CMSat::lit_Undef) continue; need += 2;
}
if (cap < need || buf == nullptr) return need;
size_t k = 0;
for (const auto &kv : m) {
if (kv.second == CMSat::lit_Undef) continue;
buf[k++] = (int32_t)(kv.first + 1u); buf[k++] = lit_to_dimacs(kv.second); }
return need;
}
size_t arjun_shim_cur_multiplier(ArjunShim *s, char *buf, size_t cap) {
std::ostringstream os;
s->cur->get_multiplier_weight()->display(os);
std::string str = os.str();
if (cap == 0 || buf == nullptr) return str.size();
size_t n = str.size() < cap - 1 ? str.size() : cap - 1;
std::memcpy(buf, str.data(), n);
buf[n] = '\0';
return str.size();
}
void arjun_shim_set_lit_weight(ArjunShim *s, int32_t lit, const char *weight_str) {
uint32_t var = (uint32_t)(lit < 0 ? -lit : lit) - 1u; Lit l(var, lit < 0);
std::unique_ptr<Field> w(s->fg->zero());
std::string tok(weight_str);
tok += " 0";
w->parse(tok, 0);
s->cur->set_lit_weight(l, w);
}
void arjun_shim_clean_sampl(ArjunShim *s) {
s->cur->start_with_clean_sampl_vars();
}
size_t arjun_shim_lit_weight(ArjunShim *s, int32_t lit, char *buf, size_t cap) {
uint32_t var = (uint32_t)(lit < 0 ? -lit : lit) - 1u;
Lit l(var, lit < 0);
std::unique_ptr<Field> w = s->cur->get_lit_weight(l);
std::ostringstream os;
w->display(os);
std::string str = os.str();
if (cap == 0 || buf == nullptr) return str.size();
size_t n = str.size() < cap - 1 ? str.size() : cap - 1;
std::memcpy(buf, str.data(), n);
buf[n] = '\0';
return str.size();
}
}