Download code/llmc_training/llmc/sampler.h from SciCode/MileTone_2: direct link, hf CLI and curl.
- Browser
- Download file 1.15 kB
-
https://huggingface.co/SciCode/MileTone_2/resolve/main/code/llmc_training/llmc/sampler.h
- Command line
-
hf download hf://SciCode/MileTone_2/code/llmc_training/llmc/sampler.h
-
curl -L -o sampler.h https://huggingface.co/SciCode/MileTone_2/resolve/main/code/llmc_training/llmc/sampler.h
1.15 kB
| /* | |
| Implements a simple Sampler, used during model inference to sample tokens. | |
| */ | |
| // Simple xorshift RNG | |
| unsigned int random_u32(unsigned long long *state) { | |
| // xorshift rng: https://en.wikipedia.org/wiki/Xorshift#xorshift.2A | |
| *state ^= *state >> 12; | |
| *state ^= *state << 25; | |
| *state ^= *state >> 27; | |
| return (*state * 0x2545F4914F6CDD1Dull) >> 32; | |
| } | |
| float random_f32(unsigned long long *state) { // random float32 in [0,1) | |
| return (random_u32(state) >> 8) / 16777216.0f; | |
| } | |
| int sample_softmax(const float* logits, int n, float coin) { | |
| // sample index from logits (converted to probabilities using softmax) | |
| // coin is a random number in [0, 1), usually from random_f32() | |
| double norm = 0; | |
| for (int i = 0; i < n; i++) { | |
| norm += expf(logits[i]); | |
| } | |
| // instead of dividing all exp(logits), we can just multiply coin. | |
| coin *= norm; | |
| float cdf = 0.0f; | |
| for (int i = 0; i < n; i++) { | |
| cdf += expf(logits[i]); | |
| if (coin < cdf) { | |
| return i; | |
| } | |
| } | |
| return n - 1; // in case of rounding errors | |
| } | |