| 235 | } |
| 236 | |
| 237 | uint32_t Model::decode(const std::vector<uint32_t>& tokens, float temperature, float top_p, |
| 238 | size_t top_k, const std::string& profile_file, float* out_entropy, |
| 239 | float min_p, float repetition_penalty) { |
| 240 | |
| 241 | if (temperature < 0) { |
| 242 | temperature = config_.default_temperature; |
| 243 | } |
| 244 | if (top_p < 0) { |
| 245 | top_p = config_.default_top_p; |
| 246 | } |
| 247 | if (top_k == 0) { |
| 248 | top_k = config_.default_top_k; |
| 249 | } |
| 250 | auto final_hidden = forward(tokens, true); |
| 251 | |
| 252 | auto* gb = static_cast<CactusGraph*>(graph_handle_); |
| 253 | auto backend = config_.default_backend == Config::Backend::CPU |
| 254 | ? ComputeBackend::CPU |
| 255 | : ComputeBackend::NPU; |
| 256 | |
| 257 | auto last_hidden = gb->index(final_hidden, tokens.size() - 1, 0); |
| 258 | const auto& last_hidden_buf = gb->get_output_buffer(last_hidden); |
| 259 | size_t hidden_dim = last_hidden_buf.shape[0]; |
| 260 | last_hidden = gb->reshape(last_hidden, {1, hidden_dim}); |
| 261 | |
| 262 | auto logits_node_id = gb->matmul(last_hidden, output_weight_node_id_, true, backend); |
| 263 | |
| 264 | if (config_.final_logit_softcapping > 0.0f) { |
| 265 | float inv_cap = 1.0f / config_.final_logit_softcapping; |
| 266 | logits_node_id = gb->scalar_multiply(logits_node_id, inv_cap); |
| 267 | logits_node_id = gb->tanh(logits_node_id); |
| 268 | logits_node_id = gb->scalar_multiply(logits_node_id, config_.final_logit_softcapping); |
| 269 | } |
| 270 | auto sampled_token_id = sample_token(gb, logits_node_id, temperature, top_p, top_k, min_p, repetition_penalty); |
| 271 | |
| 272 | gb->execute(profile_file); |
| 273 | |
| 274 | compute_entropy(gb, logits_node_id, out_entropy); |
| 275 | |
| 276 | post_execute_updates(gb, tokens.size()); |
| 277 | update_kv_cache(gb, tokens.size()); |
| 278 | |
| 279 | auto* output_ptr = gb->get_output(sampled_token_id); |
| 280 | uint32_t result_token = *static_cast<uint32_t*>(output_ptr); |
| 281 | record_sampled_token(result_token); |
| 282 | return result_token; |
| 283 | } |
| 284 | |
| 285 | size_t Model::sample_token(CactusGraph* gb, size_t logits_node_id, float temperature, float top_p, size_t top_k, |
| 286 | float min_p, float repetition_penalty, |
nothing calls this directly
no test coverage detected