(
model, tokenizer, params, device, context_len=8192, stream_interval=2, force_generate=False
)
| 55 | |
| 56 | @torch.inference_mode() |
| 57 | def generate_stream( |
| 58 | model, tokenizer, params, device, context_len=8192, stream_interval=2, force_generate=False |
| 59 | ): |
| 60 | prompt = params["prompt"] |
| 61 | len_prompt = len(prompt) |
| 62 | temperature = float(params.get("temperature", 1.0)) |
| 63 | repetition_penalty = float(params.get("repetition_penalty", 1.0)) |
| 64 | top_p = float(params.get("top_p", 1.0)) |
| 65 | top_k = int(params.get("top_k", -1)) # -1 means disable |
| 66 | max_new_tokens = int(params.get("max_new_tokens", 256)) |
| 67 | stop_str = params.get("stop", None) |
| 68 | echo = bool(params.get("echo", True)) |
| 69 | stop_token_ids = params.get("stop_token_ids", None) or [] |
| 70 | stop_token_ids.append(tokenizer.eos_token_id) |
| 71 | |
| 72 | logits_processor = prepare_logits_processor( |
| 73 | temperature, repetition_penalty, top_p, top_k |
| 74 | ) |
| 75 | |
| 76 | input_ids = tokenizer(prompt).input_ids |
| 77 | input_echo_len = len(input_ids) |
| 78 | output_ids = list(input_ids) |
| 79 | |
| 80 | if model.config.is_encoder_decoder: |
| 81 | max_src_len = context_len |
| 82 | else: |
| 83 | max_src_len = context_len - max_new_tokens - 8 |
| 84 | |
| 85 | input_ids = input_ids[-max_src_len:] |
| 86 | |
| 87 | if model.config.is_encoder_decoder: |
| 88 | encoder_output = model.encoder( |
| 89 | input_ids=torch.as_tensor([input_ids], device=device) |
| 90 | )[0] |
| 91 | start_ids = torch.as_tensor( |
| 92 | [[model.generation_config.decoder_start_token_id]], |
| 93 | dtype=torch.int64, |
| 94 | device=device, |
| 95 | ) |
| 96 | |
| 97 | past_key_values = out = None |
| 98 | for i in range(max_new_tokens): |
| 99 | if i == 0: |
| 100 | if model.config.is_encoder_decoder: |
| 101 | out = model.decoder( |
| 102 | input_ids=start_ids, |
| 103 | encoder_hidden_states=encoder_output, |
| 104 | use_cache=True, |
| 105 | ) |
| 106 | logits = model.lm_head(out[0]) |
| 107 | else: |
| 108 | out = model(torch.as_tensor([input_ids], device=device), use_cache=True) |
| 109 | logits = out.logits |
| 110 | past_key_values = out.past_key_values |
| 111 | else: |
| 112 | if model.config.is_encoder_decoder: |
| 113 | out = model.decoder( |
| 114 | input_ids=torch.as_tensor([[token]], device=device), |
nothing calls this directly
no test coverage detected