| 68 | } |
| 69 | |
| 70 | void GetSegmentPredictions( |
| 71 | const std::vector<std::string>& input, |
| 72 | const ::tflite::FlatBufferModel& model, const SmartReplyConfig& config, |
| 73 | std::vector<PredictorResponse>* predictor_responses) { |
| 74 | // Initialize interpreter |
| 75 | std::unique_ptr<::tflite::Interpreter> interpreter; |
| 76 | ::tflite::MutableOpResolver resolver; |
| 77 | RegisterSelectedOps(&resolver); |
| 78 | ::tflite::InterpreterBuilder(model, resolver)(&interpreter); |
| 79 | |
| 80 | if (!model.initialized()) { |
| 81 | fprintf(stderr, "Failed to mmap model \n"); |
| 82 | return; |
| 83 | } |
| 84 | |
| 85 | // Execute Tflite Model |
| 86 | std::map<std::string, float> response_map; |
| 87 | std::vector<std::string> sentences; |
| 88 | for (const std::string& str : input) { |
| 89 | std::vector<std::string> splitted_str = SplitSentence(str); |
| 90 | sentences.insert(sentences.end(), splitted_str.begin(), splitted_str.end()); |
| 91 | } |
| 92 | for (const auto& sentence : sentences) { |
| 93 | ExecuteTfLite(sentence, interpreter.get(), &response_map); |
| 94 | } |
| 95 | |
| 96 | // Generate the result. |
| 97 | for (const auto& iter : response_map) { |
| 98 | PredictorResponse prediction(iter.first, iter.second); |
| 99 | predictor_responses->emplace_back(prediction); |
| 100 | } |
| 101 | std::sort(predictor_responses->begin(), predictor_responses->end(), |
| 102 | [](const PredictorResponse& a, const PredictorResponse& b) { |
| 103 | return a.GetScore() > b.GetScore(); |
| 104 | }); |
| 105 | |
| 106 | // Add backoff response. |
| 107 | for (const auto& backoff : config.backoff_responses) { |
| 108 | if (predictor_responses->size() >= config.num_response) { |
| 109 | break; |
| 110 | } |
| 111 | predictor_responses->emplace_back(backoff, config.backoff_confidence); |
| 112 | } |
| 113 | } |
| 114 | |
| 115 | } // namespace smartreply |
| 116 | } // namespace custom |