| 6 | namespace vlm { |
| 7 | |
| 8 | mx::array topPSampling(const mx::array &Logits, float TopP, float Temperature) { |
| 9 | |
| 10 | mx::array WorkingLogits = Logits; |
| 11 | |
| 12 | if (WorkingLogits.dtype() == mx::bfloat16) { |
| 13 | WorkingLogits = astype(WorkingLogits, mx::float32); |
| 14 | } |
| 15 | |
| 16 | mx::array Probs = mx::softmax(WorkingLogits / Temperature, -1); |
| 17 | mx::array SortedIndices = mx::argsort(Probs, -1); |
| 18 | mx::array SqueezedIndices = mx::squeeze(SortedIndices, 0); |
| 19 | mx::array SortedProbs = mx::take(Probs, SqueezedIndices, -1); |
| 20 | mx::array CumulativeProbs = mx::cumsum(SortedProbs, -1); |
| 21 | mx::array TopProbs = mx::where(CumulativeProbs > 1.0f - TopP, SortedProbs, |
| 22 | mx::zeros_like(SortedProbs)); |
| 23 | mx::array SortedToken = mx::random::categorical(mx::log(TopProbs)); |
| 24 | mx::array Token = mx::take(SqueezedIndices, SortedToken); |
| 25 | |
| 26 | return Token; |
| 27 | } |
| 28 | |
| 29 | } // namespace vlm |
| 30 | } // namespace WasmEdge::Host::WASINN::MLX |