MCPcopy Create free account
hub / github.com/0xShug0/audio.cpp / PrefillGraph

Method PrefillGraph

src/models/ace_step/planner.cpp:1756–1818  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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.

Callers

nothing calls this directly

Calls 15

ggml_initFunction · 0.85
ggml_new_tensor_1dFunction · 0.85
wrap_tensorFunction · 0.85
EmbeddingModuleClass · 0.85
reshape_tensorFunction · 0.85
ggml_new_tensor_4dFunction · 0.85
SliceModuleClass · 0.85
RMSNormModuleClass · 0.85
LinearModuleClass · 0.85
ggml_set_outputFunction · 0.85
ggml_new_graph_customFunction · 0.85

Tested by

no test coverage detected