| 104 | } |
| 105 | |
| 106 | std::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 | |
| 150 | HeartMuLaFrameEmbeddingInputs prompt_embedding_inputs( |
| 151 | const HeartMuLaPromptEncoding & encoding, |
no test coverage detected