25 using SolverFactory = std::function<std::unique_ptr<MaxFlowSolver<T>>(
int num_vars,
int num_edges)>;
31 : model_(model), solver_factory_(std::move(solver_factory)) {
49 const std::vector<int> active_nodes = model_.get_active_nodes(alpha_label);
50 if (active_nodes.empty())
return false;
52 std::vector<typename MaxFlowSolver<T>::Var> node_var_ids;
53 T old_subgraph_energy = 0;
54 const auto solver = build_expansion_graph(alpha_label, active_nodes, node_var_ids, old_subgraph_energy);
56 const T new_subgraph_energy = solver->minimize();
59 if constexpr (std::is_floating_point_v<T>) {
60 improved = (old_subgraph_energy - new_subgraph_energy >
static_cast<T
>(1e-5));
62 improved = (new_subgraph_energy < old_subgraph_energy);
64 if (!improved)
return false;
67 for (
const int node: active_nodes) {
69 model_.set_label(node, alpha_label);
80 std::unique_ptr<MaxFlowSolver<T>> build_expansion_graph(
const int alpha_label,
const std::vector<int> &active_nodes,
82 T &old_subgraph_energy)
const {
83 const int num_active = active_nodes.size();
84 if (num_active == 0)
return nullptr;
86 const int estimated_edges = num_active * 4;
87 auto solver = solver_factory_(num_active, estimated_edges);
89 node_var_ids.assign(model_.num_nodes(), -1);
91 for (
int i = 0; i < num_active; ++i) {
92 int node = active_nodes[i];
93 node_var_ids[node] = solver->add_variable();
98 for (
const int node: active_nodes) {
100 const int current_label = model_.get_label(node);
101 const T e0 = model_.get_unary_cost(node, alpha_label);
102 const T e1 = model_.get_unary_cost(node, current_label);
103 solver->add_term1(var, e0, e1);
107 for (
const int node_i: active_nodes) {
108 const int current_label_i = model_.get_label(node_i);
111 for (
const int node_j: model_.get_neighbors_range(node_i)) {
112 if (node_var_ids[node_j] != -1) {
113 if (node_i < node_j) {
115 const int current_label_j = model_.get_label(node_j);
116 const T e00 = model_.get_pairwise_cost(node_i, node_j, alpha_label, alpha_label);
117 const T e01 = model_.get_pairwise_cost(node_i, node_j, alpha_label, current_label_j);
118 const T e10 = model_.get_pairwise_cost(node_i, node_j, current_label_i, alpha_label);
119 const T e11 = model_.get_pairwise_cost(node_i, node_j, current_label_i, current_label_j);
120 solver->add_term2(var_i, var_j, e00, e01, e10, e11);
124 const int current_label_j = model_.get_label(node_j);
125 const T e0 = model_.get_pairwise_cost(node_i, node_j, alpha_label, current_label_j);
126 const T e1 = model_.get_pairwise_cost(node_i, node_j, current_label_i, current_label_j);
127 solver->add_term1(var_i, e0, e1);
133 old_subgraph_energy = old_energy;
139 std::function<void()> on_iteration_ = []{};