MCPcopy Create free account
hub / github.com/OpenBMB/ToolBench / generate_stream

Function generate_stream

toolbench/inference/utils.py:57–220  ·  view source on GitHub ↗
(
    model, tokenizer, params, device, context_len=8192, stream_interval=2, force_generate=False
)

Source from the content-addressed store, hash-verified

55
56@torch.inference_mode()
57def 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),

Callers

nothing calls this directly

Calls 1

prepare_logits_processorFunction · 0.85

Tested by

no test coverage detected