(log_probs_revised, top_p, top_k_num, use_pynative=False, bad_words_index=[])
| 32 | |
| 33 | |
| 34 | def sampler(log_probs_revised, top_p, top_k_num, use_pynative=False, bad_words_index=[]): |
| 35 | for i, bad_words in enumerate(bad_words_index): |
| 36 | for bad_word in bad_words: |
| 37 | log_probs_revised[i, bad_word] = -10000 |
| 38 | """Convert the log_probs to probability""" |
| 39 | if use_pynative: |
| 40 | log_probs_revised = log_probs_revised.asnumpy() |
| 41 | |
| 42 | return log_probs_revised.argmax(axis=1) |
| 43 | |
| 44 | |
| 45 | def generate_increment(model, origin_inputs, origin_length, config, tokenizer, verbose=False): |
no outgoing calls
no test coverage detected