| 19 | std::unique_ptr<Sampling> NewPartialSumSampling(const vec_float_t* probs); |
| 20 | |
| 21 | std::unique_ptr<Sampling> NewSampling(const vec_float_t* probs, |
| 22 | SamplingEnum type) { |
| 23 | std::unique_ptr<Sampling> sampling; |
| 24 | switch (type) { |
| 25 | case SamplingEnum::UNIFORM: |
| 26 | sampling = NewUniformSampling(probs); |
| 27 | break; |
| 28 | case SamplingEnum::ALIAS: |
| 29 | sampling = NewAliasSampling(probs); |
| 30 | break; |
| 31 | case SamplingEnum::WORD2VEC: |
| 32 | sampling = NewWord2vecSampling(probs); |
| 33 | break; |
| 34 | case SamplingEnum::PARTIAL_SUM: |
| 35 | sampling = NewPartialSumSampling(probs); |
| 36 | break; |
| 37 | default: |
| 38 | DXERROR( |
| 39 | "Need type: UNIFORM(0) || ALIAS(1) || WORD2VEC(2) || PARTIAL_SUM(3), " |
| 40 | "got type: %d.", |
| 41 | (int)type); |
| 42 | break; |
| 43 | } |
| 44 | return sampling; |
| 45 | } |
| 46 | |
| 47 | } // namespace embedx |