| 129 | } |
| 130 | |
| 131 | MioTTSGenerationOptions generation_options_from_request( |
| 132 | const MioTTSConfig & config, |
| 133 | const runtime::TaskRequest & request) { |
| 134 | MioTTSGenerationOptions out; |
| 135 | out.max_tokens = config.max_tokens; |
| 136 | out.top_k = config.top_k; |
| 137 | out.top_p = config.top_p; |
| 138 | out.temperature = config.temperature; |
| 139 | out.repetition_penalty = config.repetition_penalty; |
| 140 | out.presence_penalty = config.presence_penalty; |
| 141 | out.frequency_penalty = config.frequency_penalty; |
| 142 | out.do_sample = config.do_sample; |
| 143 | if (const auto value = runtime::parse_int_option(request.options, {"max_tokens"})) { |
| 144 | out.max_tokens = *value; |
| 145 | } |
| 146 | if (const auto value = runtime::parse_int_option(request.options, {"top_k"})) { |
| 147 | out.top_k = *value; |
| 148 | } |
| 149 | if (const auto value = runtime::parse_float_option(request.options, {"top_p"})) { |
| 150 | out.top_p = *value; |
| 151 | } |
| 152 | if (const auto value = runtime::parse_float_option(request.options, {"temperature"})) { |
| 153 | out.temperature = *value; |
| 154 | } |
| 155 | if (const auto value = runtime::parse_float_option(request.options, {"repetition_penalty"})) { |
| 156 | out.repetition_penalty = *value; |
| 157 | } |
| 158 | if (const auto value = runtime::parse_float_option(request.options, {"presence_penalty"})) { |
| 159 | out.presence_penalty = *value; |
| 160 | } |
| 161 | if (const auto value = runtime::parse_float_option(request.options, {"frequency_penalty"})) { |
| 162 | out.frequency_penalty = *value; |
| 163 | } |
| 164 | out.seed = runtime::parse_u32_option(request.options, {"seed"}) |
| 165 | .value_or(runtime::random_u32_seed()); |
| 166 | if (const auto value = runtime::find_option(request.options, {"do_sample"})) { |
| 167 | out.do_sample = runtime::parse_bool_option(*value, "do_sample"); |
| 168 | } |
| 169 | if (out.max_tokens <= 0) { |
| 170 | throw std::runtime_error("MioTTS max_tokens must be positive"); |
| 171 | } |
| 172 | if (out.temperature < 0.0F || out.temperature > 2.0F) { |
| 173 | throw std::runtime_error("MioTTS temperature must be in [0, 2]"); |
| 174 | } |
| 175 | if (out.top_p < 0.0F || out.top_p > 1.0F) { |
| 176 | throw std::runtime_error("MioTTS top_p must be in [0, 1]"); |
| 177 | } |
| 178 | if (out.repetition_penalty < 1.0F || out.repetition_penalty > 1.5F) { |
| 179 | throw std::runtime_error("MioTTS repetition_penalty must be in [1, 1.5]"); |
| 180 | } |
| 181 | if (out.presence_penalty < 0.0F || out.presence_penalty > 1.0F) { |
| 182 | throw std::runtime_error("MioTTS presence_penalty must be in [0, 1]"); |
| 183 | } |
| 184 | if (out.frequency_penalty < 0.0F || out.frequency_penalty > 1.0F) { |
| 185 | throw std::runtime_error("MioTTS frequency_penalty must be in [0, 1]"); |
| 186 | } |
| 187 | return out; |
| 188 | } |
no test coverage detected