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

Function prompt_embedding_inputs

src/models/heartmula/generator.cpp:150–196  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

148}
149
150HeartMuLaFrameEmbeddingInputs prompt_embedding_inputs(
151 const HeartMuLaPromptEncoding & encoding,
152 const HeartMuLaConfig & config,
153 bool use_cfg) {
154 HeartMuLaFrameEmbeddingInputs inputs;
155 inputs.batch_size = encoding.batch_size;
156 inputs.steps = encoding.prompt_len;
157 const size_t batch = static_cast<size_t>(encoding.batch_size);
158 const size_t steps = static_cast<size_t>(encoding.prompt_len);
159 const size_t lanes = static_cast<size_t>(encoding.parallel_number);
160 const size_t codebooks = static_cast<size_t>(config.audio_num_codebooks);
161 inputs.audio_token_ids.assign(batch * steps * codebooks, 0);
162 inputs.text_token_ids.assign(batch * steps, 0);
163 inputs.audio_mask.assign(batch * steps * codebooks, 0.0F);
164 inputs.text_cond_mask.assign(batch * steps, 0.0F);
165 inputs.text_uncond_mask.assign(batch * steps, 0.0F);
166 const size_t actual_batch = use_cfg ? batch / 2 : batch;
167 for (size_t b = 0; b < batch; ++b) {
168 const bool uncond = use_cfg && b >= actual_batch;
169 for (size_t t = 0; t < steps; ++t) {
170 for (size_t q = 0; q < codebooks; ++q) {
171 const size_t src = prompt_flat(b, t, q, steps, lanes);
172 const size_t dst = (b * steps + t) * codebooks + q;
173 inputs.audio_token_ids[dst] = static_cast<int32_t>(
174 encoding.tokens[src] + static_cast<int64_t>(q) * config.audio_vocab_size);
175 inputs.audio_mask[dst] = encoding.tokens_mask[src] != 0U ? 1.0F : 0.0F;
176 }
177 const size_t text_src = prompt_flat(b, t, lanes - 1, steps, lanes);
178 inputs.text_token_ids[b * steps + t] = static_cast<int32_t>(encoding.tokens[text_src]);
179 const float text_mask = encoding.tokens_mask[text_src] != 0U ? 1.0F : 0.0F;
180 inputs.text_cond_mask[b * steps + t] = uncond ? 0.0F : text_mask;
181 inputs.text_uncond_mask[b * steps + t] = uncond ? text_mask : 0.0F;
182 }
183 }
184 inputs.apply_muq = true;
185 inputs.muq_row = encoding.muq_idx.empty() ? 0 : encoding.muq_idx.front();
186 inputs.muq_embed = encoding.muq_embed;
187 inputs.muq_cond_mask.assign(batch, 1.0F);
188 inputs.muq_uncond_mask.assign(batch, 0.0F);
189 if (use_cfg) {
190 for (size_t b = actual_batch; b < batch; ++b) {
191 inputs.muq_cond_mask[b] = 0.0F;
192 inputs.muq_uncond_mask[b] = 1.0F;
193 }
194 }
195 return inputs;
196}
197
198HeartMuLaFrameEmbeddingInputs audio_frame_embedding_inputs(
199 const std::vector<int32_t> & frame_tokens,

Callers 1

Calls 3

prompt_flatFunction · 0.85
assignMethod · 0.80
emptyMethod · 0.45

Tested by

no test coverage detected