| 22 | |
| 23 | template <typename OptimizableOBJ> |
| 24 | double |
| 25 | Optimizer<OptimizableOBJ>::Optimize(std::shared_ptr<OptimizableOBJ> obj) { |
| 26 | vsag::logger::info(fmt::format("============Start Optimize===========")); |
| 27 | bool have_improvement = false; |
| 28 | SearchStatistics stats; |
| 29 | double original_loss = obj->MockRun(stats); |
| 30 | |
| 31 | // generate a group of runtime params |
| 32 | UnorderedMap<std::string, float> current_params(allocator_); |
| 33 | for (auto& param : parameters_) { |
| 34 | bool successful_optimized = false; |
| 35 | double before_loss = obj->MockRun(stats); |
| 36 | double best_loss = before_loss; |
| 37 | double best_improve = 0; |
| 38 | |
| 39 | while (not param.IsEnd()) { |
| 40 | current_params[param.name_] = param.Cur(); |
| 41 | auto set_status = obj->SetRuntimeParameters(current_params); |
| 42 | if (not set_status) { |
| 43 | continue; |
| 44 | } |
| 45 | |
| 46 | // evaluate |
| 47 | double loss = obj->MockRun(stats); |
| 48 | double improvement = (before_loss - loss) / before_loss * 100; |
| 49 | vsag::logger::info( |
| 50 | fmt::format("setting {} -> {}, get new loss = {:.3f} from before = {:.3f}, " |
| 51 | "improving = {:.3f} %", |
| 52 | param.name_, |
| 53 | current_params[param.name_], |
| 54 | loss, |
| 55 | before_loss, |
| 56 | improvement)); |
| 57 | |
| 58 | // update |
| 59 | // overall principle: choose smaller param |
| 60 | if (loss < best_loss and // condition1. has improvement |
| 61 | improvement - best_improve > 0.2 and // condition2. improvement is noteworthy |
| 62 | improvement > 2.0) { // condition3. improvement is valid |
| 63 | successful_optimized = true; |
| 64 | have_improvement = true; |
| 65 | best_loss = loss; |
| 66 | best_improve = improvement; |
| 67 | best_params_[param.name_] = current_params[param.name_]; |
| 68 | } |
| 69 | |
| 70 | param.Next(); |
| 71 | } |
| 72 | |
| 73 | if (successful_optimized) { |
| 74 | current_params[param.name_] = best_params_[param.name_]; |
| 75 | vsag::logger::info(fmt::format("setting to best param: {} -> {}, improving {:.3f}%", |
| 76 | param.name_, |
| 77 | best_params_[param.name_], |
| 78 | best_improve)); |
| 79 | } else { |
| 80 | param.Reset(); |
| 81 | current_params[param.name_] = param.Cur(); |