use std::time::{Duration, Instant};
use crate::hash::FxHashMap;
use crate::packed::PackedTask;
use crate::pddl3::plan_cost;
use crate::search::{solve_subgoal_bounded, SatGuidance, SearchCfg};
pub struct EspcResult {
pub ops: Vec<usize>,
pub cost: f64,
pub iterations: usize,
}
fn env_i64(key: &str, default: i64) -> i64 {
std::env::var(key)
.ok()
.and_then(|s| s.parse::<i64>().ok())
.unwrap_or(default)
}
fn measure_violation(
task: &PackedTask,
ops: &[usize],
base: &[(u32, u32, i64)],
) -> FxHashMap<u32, i64> {
let mut s = task.initial();
for &oi in ops {
s = task.apply(oi, &s);
}
let mut v: FxHashMap<u32, i64> = FxHashMap::default();
for &(m, d, _) in base {
if crate::bitset::test(&s.bits, m as usize) && !crate::bitset::test(&s.bits, d as usize) {
*v.entry(m).or_insert(0) += 1;
}
}
v
}
fn solve_under_penalties(
task: &PackedTask,
cost_fluent: usize,
sat: &SatGuidance,
threads: usize,
cfg: SearchCfg,
deadline: Instant,
) -> Option<(Vec<usize>, f64)> {
const INNER_MAX: usize = 200;
let init = task.initial();
let mut bound = f64::INFINITY;
let mut best: Option<(Vec<usize>, f64)> = None;
for i in 0..INNER_MAX {
if i > 0 && Instant::now() >= deadline {
break;
}
let (opt, _capped) = solve_subgoal_bounded(
task,
&init,
&task.goal_pos,
&task.goal_num,
cost_fluent,
bound,
threads,
cfg,
Some(sat),
);
match opt {
Some(ops) => {
let cost = plan_cost(task, &ops, cost_fluent);
best = Some((ops, cost));
if cost <= 0.0 {
break;
}
bound = cost; }
None => break,
}
}
best
}
pub fn espc_optimize(
task: &PackedTask,
cost_fluent: usize,
sat: &mut SatGuidance,
seed: Option<(Vec<usize>, f64)>,
threads: usize,
cfg: SearchCfg,
) -> Option<EspcResult> {
let base: Vec<(u32, u32, i64)> = sat.deadline.clone();
if base.is_empty() {
return seed.map(|(ops, cost)| EspcResult {
ops,
cost,
iterations: 0,
});
}
let lambda0 = env_i64("FF_ESPC_LAMBDA0", 0).max(0);
let rate0 = env_i64("FF_ESPC_RATE", 20).max(1);
let outer_max = env_i64("FF_ESPC_OUTER", 16).max(1) as usize;
let k_bump = env_i64("FF_ESPC_K", 2).max(1); let stall_max = env_i64("FF_ESPC_STALL", 4).max(1) as usize;
let time_ms = env_i64("FF_ESPC_TIME_MS", 15_000).max(0) as u64;
let debug = std::env::var("FF_RES_DEBUG").is_ok();
let mut lambda: FxHashMap<u32, i64> = FxHashMap::default();
for &(m, _, _) in &base {
lambda.entry(m).or_insert(lambda0);
}
sat.deadline_weight = 1;
let mut best = seed;
let mut rate = rate0;
let mut consec = 0i64;
let mut stall = 0usize;
let mut iterations = 0usize;
let deadline = Instant::now() + Duration::from_millis(time_ms);
for outer in 0..outer_max {
if outer > 0 && Instant::now() >= deadline {
break;
}
iterations += 1;
for (pair, b) in sat.deadline.iter_mut().zip(&base) {
let lam = *lambda.get(&b.0).unwrap_or(&0);
pair.2 = lam.saturating_mul(b.2);
}
let Some((ops, cost)) =
solve_under_penalties(task, cost_fluent, sat, threads, cfg, deadline)
else {
break; };
let viol = measure_violation(task, &ops, &base);
let total_v: i64 = viol.values().sum();
let improved = best.as_ref().map_or(true, |(_, c)| cost < *c - 1e-9);
if improved {
best = Some((ops, cost));
stall = 0;
consec = 0;
} else {
stall += 1;
consec += 1;
}
if debug {
eprintln!(
"[ESPC] iter {outer}: cost={cost} violations={total_v} best={} rate={rate}",
best.as_ref().map(|(_, c)| *c).unwrap_or(f64::INFINITY)
);
}
if total_v == 0 {
break; }
if stall >= stall_max {
break; }
for (&m, &v) in &viol {
if v > 0 {
let e = lambda.entry(m).or_insert(lambda0);
*e = e.saturating_add(rate.saturating_mul(v));
}
}
if consec >= k_bump {
rate = rate.saturating_mul(2);
consec = 0;
}
}
best.map(|(ops, cost)| EspcResult {
ops,
cost,
iterations,
})
}