| 2261 | return runtime_.get() == &runtime && prompt_steps_ == prompt_steps; |
| 2262 | } |
| 2263 | |
| 2264 | CfgPrefillOutput run( |
| 2265 | const AceStepTokenizedText & conditional_prompt, |
| 2266 | const AceStepTokenizedText & unconditional_prompt) { |
| 2267 | const auto & config = runtime_->assets().config.planner; |
| 2268 | if (static_cast<int64_t>(conditional_prompt.input_ids.size()) != prompt_steps_ || |
| 2269 | static_cast<int64_t>(unconditional_prompt.input_ids.size()) != prompt_steps_ || |
| 2270 | static_cast<int64_t>(conditional_prompt.attention_mask.size()) != prompt_steps_ || |
| 2271 | static_cast<int64_t>(unconditional_prompt.attention_mask.size()) != prompt_steps_) { |
| 2272 | throw std::runtime_error("ACE-Step planner CFG prefill prompt size mismatch"); |
| 2273 | } |
| 2274 | std::vector<int32_t> token_ids(static_cast<size_t>(2 * prompt_steps_), 0); |
| 2275 | std::copy(conditional_prompt.input_ids.begin(), conditional_prompt.input_ids.end(), token_ids.begin()); |
| 2276 | std::copy( |
| 2277 | unconditional_prompt.input_ids.begin(), |
| 2278 | unconditional_prompt.input_ids.end(), |
| 2279 | token_ids.begin() + static_cast<std::ptrdiff_t>(prompt_steps_)); |
| 2280 | ggml_backend_tensor_set(token_ids_, token_ids.data(), 0, token_ids.size() * sizeof(int32_t)); |
| 2281 | |
| 2282 | std::vector<int32_t> position_ids(static_cast<size_t>(2 * prompt_steps_), 0); |
| 2283 | for (int64_t i = 0; i < prompt_steps_; ++i) { |
| 2284 | const int32_t position = static_cast<int32_t>(i); |
| 2285 | position_ids[static_cast<size_t>(i)] = position; |
| 2286 | position_ids[static_cast<size_t>(prompt_steps_ + i)] = position; |
| 2287 | } |
| 2288 | ggml_backend_tensor_set(positions_, position_ids.data(), 0, position_ids.size() * sizeof(int32_t)); |
| 2289 | const auto attention_mask_values = build_cfg_prefill_attention_mask_values( |
| 2290 | conditional_prompt.attention_mask, |
| 2291 | unconditional_prompt.attention_mask); |
| 2292 | ggml_backend_tensor_set( |
| 2293 | attention_mask_, |
| 2294 | attention_mask_values.data(), |
| 2295 | 0, |
| 2296 | attention_mask_values.size() * sizeof(ggml_fp16_t)); |
| 2297 | std::vector<float> query_mask_values(static_cast<size_t>(2 * prompt_steps_), 0.0F); |
| 2298 | for (int64_t i = 0; i < prompt_steps_; ++i) { |
| 2299 | query_mask_values[static_cast<size_t>(i)] = |
| 2300 | conditional_prompt.attention_mask[static_cast<size_t>(i)] != 0 ? 1.0F : 0.0F; |
| 2301 | query_mask_values[static_cast<size_t>(prompt_steps_ + i)] = |
| 2302 | unconditional_prompt.attention_mask[static_cast<size_t>(i)] != 0 ? 1.0F : 0.0F; |
| 2303 | } |
| 2304 | ggml_backend_tensor_set(query_mask_, query_mask_values.data(), 0, query_mask_values.size() * sizeof(float)); |
| 2305 | |
| 2306 | core::set_backend_threads(runtime_->backend(), runtime_->threads()); |
| 2307 | const ggml_status status = engine::core::compute_backend_graph(runtime_->backend(), graph_); |
| 2308 | ggml_backend_synchronize(runtime_->backend()); |
| 2309 | if (status != GGML_STATUS_SUCCESS) { |
| 2310 | throw std::runtime_error("ACE-Step planner CFG prefill graph compute failed"); |
| 2311 | } |
| 2312 | |
| 2313 | CfgPrefillOutput out; |
| 2314 | out.current_end = prompt_steps_; |
| 2315 | out.valid_steps = prompt_steps_; |
| 2316 | out.conditional_logits.resize(static_cast<size_t>(config.vocab_size)); |
| 2317 | out.unconditional_logits.resize(static_cast<size_t>(config.vocab_size)); |
| 2318 | ggml_backend_tensor_get( |
| 2319 | logits_cond_, |
| 2320 | out.conditional_logits.data(), |
nothing calls this directly
no test coverage detected