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

Function sample_with_cfg

src/models/heartmula/generator.cpp:106–148  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

104}
105
106std::vector<int32_t> sample_with_cfg(
107 const std::vector<float> & logits,
108 int64_t batch,
109 int64_t vocab_size,
110 float cfg_scale,
111 int64_t topk,
112 float temperature,
113 uint64_t seed,
114 uint64_t & sample_call_index,
115 const TorchCudaSamplingPolicy & policy,
116 TopKSamplerScratch & scratch) {
117 if (batch <= 0 || vocab_size <= 0 || static_cast<int64_t>(logits.size()) != batch * vocab_size) {
118 throw std::runtime_error("HeartMuLa sampler logits shape mismatch");
119 }
120 const bool use_cfg = cfg_scale > 1.0F && batch > 1 && batch % 2 == 0;
121 const int64_t rows = use_cfg ? batch / 2 : batch;
122 std::vector<int32_t> out(static_cast<size_t>(batch), 0);
123 std::vector<float> guided(static_cast<size_t>(vocab_size));
124 for (int64_t row = 0; row < rows; ++row) {
125 const float * source = logits.data() + static_cast<size_t>(row * vocab_size);
126 if (use_cfg) {
127 const float * uncond = logits.data() + static_cast<size_t>((row + rows) * vocab_size);
128 for (int64_t v = 0; v < vocab_size; ++v) {
129 guided[static_cast<size_t>(v)] = uncond[v] + (source[v] - uncond[v]) * cfg_scale;
130 }
131 source = guided.data();
132 }
133 const int32_t token = sample_topk_row(
134 source,
135 vocab_size,
136 topk,
137 temperature,
138 seed,
139 sample_call_index++,
140 policy,
141 scratch);
142 out[static_cast<size_t>(row)] = token;
143 if (use_cfg) {
144 out[static_cast<size_t>(row + rows)] = token;
145 }
146 }
147 return out;
148}
149
150HeartMuLaFrameEmbeddingInputs prompt_embedding_inputs(
151 const HeartMuLaPromptEncoding & encoding,

Callers 1

Calls 3

sample_topk_rowFunction · 0.85
sizeMethod · 0.45
dataMethod · 0.45

Tested by

no test coverage detected