#include "objective.h"
#include <symengine/add.h>
#include <symengine/mul.h>
#include <symengine/symbol.h>
#include <symengine/integer.h>
#include <stdexcept>
using namespace SymEngine;
CompiledObjective::CompiledObjective(const Topology &t, const std::string &kind, int mode,
                                     double duty, const std::vector<double> &weights) {
    const unsigned n = kind == "ssl" ? t.N_caps : t.N_sw;
    std::vector<RCP<const Basic>> x;
    for (unsigned i=0; i<n; ++i) x.push_back(symbol("objective_x" + std::to_string(i)));
    map_basic_basic fixed;
    fixed[symbol("D")] = real_double(duty); fixed[t.frequency] = real_double(1);
    for (auto &e : t.capacitor_esr) fixed[e] = real_double(0);
    for (unsigned i=0; i<t.capacitances.size(); ++i)
        fixed[t.capacitances[i]] = kind == "ssl" ? RCP<const Basic>(x[i % n]) : RCP<const Basic>(real_double(1));
    for (unsigned i=0; i<t.switch_resistances.size(); ++i)
        fixed[t.switch_resistances[i]] = kind == "fsl" ? RCP<const Basic>(div(integer(1), x[i % n])) : RCP<const Basic>(real_double(1));
    const auto &z = kind == "ssl" ? t.ZSSL : t.ZFSL;
    RCP<const Basic> sum = integer(0);
    for (unsigned i=0; i<z.nrows(); ++i) for (unsigned j=0; j<z.ncols(); ++j) {
        auto term = z.get(i,j)->subs(fixed);
        if (mode == -2) term = div(term, mul(t.m_ratios.get(i,0), t.m_ratios.get(j,0))->subs(fixed));
        if (mode == -3) term = mul(term, real_double(weights[i] * weights[j]));
        sum = add(sum, term);
    }
    auto free = free_symbols(*sum);
    if (free.size() != n) throw std::invalid_argument("compiled objective retained unexpected symbols");
    for (auto &v : x) if (!free.count(v)) throw std::invalid_argument("compiled objective missing allocation symbol");
    lambda_.init(x, *sum, false);
}
double CompiledObjective::call(const std::vector<double> &values) { return lambda_.call(values); }
