| 1754 | size_t graph_arena_bytes) |
| 1755 | : runtime_(std::move(runtime)), |
| 1756 | prompt_steps_(prompt_steps) { |
| 1757 | if (prompt_steps_ <= 0) { |
| 1758 | throw std::runtime_error("ACE-Step planner prefill requires positive prompt length"); |
| 1759 | } |
| 1760 | ggml_init_params params{graph_arena_bytes, nullptr, true}; |
| 1761 | ctx_.reset(ggml_init(params)); |
| 1762 | if (ctx_ == nullptr) { |
| 1763 | throw std::runtime_error("failed to initialize ACE-Step planner prefill graph context"); |
| 1764 | } |
| 1765 | const auto & config = runtime_->assets().config.planner; |
| 1766 | const auto & weights = runtime_->weights(); |
| 1767 | core::ModuleBuildContext ctx{ctx_.get(), "ace_step.planner.prefill", runtime_->backend_type()}; |
| 1768 | token_ids_ = ggml_new_tensor_1d(ctx_.get(), GGML_TYPE_I32, prompt_steps_); |
| 1769 | auto token_ids = core::wrap_tensor(token_ids_, core::TensorShape::from_dims({prompt_steps_}), GGML_TYPE_I32); |
| 1770 | auto x = modules::EmbeddingModule({config.vocab_size, config.hidden_size}) |
| 1771 | .build(ctx, token_ids, weights.token_embedding); |
| 1772 | x = core::reshape_tensor(ctx, x, core::TensorShape::from_dims({1, prompt_steps_, config.hidden_size})); |
| 1773 | positions_ = ggml_new_tensor_1d(ctx_.get(), GGML_TYPE_I32, prompt_steps_); |
| 1774 | auto positions = core::wrap_tensor(positions_, core::TensorShape::from_dims({prompt_steps_}), GGML_TYPE_I32); |
| 1775 | attention_mask_ = ggml_new_tensor_4d(ctx_.get(), GGML_TYPE_F16, prompt_steps_, prompt_steps_, 1, 1); |
| 1776 | auto attention_mask = core::wrap_tensor( |
| 1777 | attention_mask_, |
| 1778 | core::TensorShape::from_dims({1, 1, prompt_steps_, prompt_steps_}), |
| 1779 | GGML_TYPE_F16); |
| 1780 | query_mask_ = ggml_new_tensor_4d(ctx_.get(), GGML_TYPE_F32, 1, 1, prompt_steps_, 1); |
| 1781 | auto query_mask = core::wrap_tensor( |
| 1782 | query_mask_, |
| 1783 | core::TensorShape::from_dims({1, prompt_steps_, 1, 1}), |
| 1784 | GGML_TYPE_F32); |
| 1785 | for (size_t layer_index = 0; layer_index < weights.layers.layers.size(); ++layer_index) { |
| 1786 | const auto & layer = weights.layers.layers[layer_index]; |
| 1787 | auto out = planner_decoder_layer_batched( |
| 1788 | ctx, |
| 1789 | x, |
| 1790 | positions, |
| 1791 | layer, |
| 1792 | config, |
| 1793 | attention_mask, |
| 1794 | query_mask, |
| 1795 | GGML_TYPE_F32); |
| 1796 | x = out.output; |
| 1797 | keys_.push_back(out.key.tensor); |
| 1798 | values_.push_back(out.value.tensor); |
| 1799 | } |
| 1800 | x = modules::SliceModule({1, prompt_steps_ - 1, 1}).build(ctx, x); |
| 1801 | x = modules::RMSNormModule({config.hidden_size, config.rms_norm_eps, true, false}) |
| 1802 | .build(ctx, x, {weights.norm, std::nullopt}); |
| 1803 | auto logits = modules::LinearModule({config.hidden_size, config.vocab_size, false}) |
| 1804 | .build(ctx, x, {weights.lm_head, std::nullopt}); |
| 1805 | logits_ = logits.tensor; |
| 1806 | ggml_set_output(logits_); |
| 1807 | graph_ = ggml_new_graph_custom(ctx_.get(), 65536, false); |
| 1808 | ggml_build_forward_expand(graph_, logits_); |
| 1809 | for (auto * input : {token_ids_, positions_, attention_mask_, query_mask_}) { |
| 1810 | ggml_set_input(input); |
| 1811 | } |
| 1812 | // The CPU reads every layer's KV outputs after graph execution. They |
| 1813 | // must remain live even after the attention nodes have consumed them. |
nothing calls this directly
no test coverage detected