MCPcopy Create free account
hub / github.com/WasmEdge/WasmEdge / topPSampling

Function topPSampling

plugins/wasi_nn/MLX/model/vlm_sampling.cpp:8–27  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

6namespace vlm {
7
8mx::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

Callers 1

generateMethod · 0.85

Calls 1

takeFunction · 0.85

Tested by

no test coverage detected