| 118 | |
| 119 | |
| 120 | class VllmDecoder(DecoderBase): |
| 121 | def __init__(self, name: str, dataset: str, tp: int, **kwargs) -> None: |
| 122 | super().__init__(name, **kwargs) |
| 123 | |
| 124 | kwargs = { |
| 125 | "tensor_parallel_size": int(os.getenv("VLLM_N_GPUS", tp)), |
| 126 | "dtype": self.dtype, |
| 127 | "trust_remote_code": self.trust_remote_code, |
| 128 | } |
| 129 | |
| 130 | self.tokenizer = AutoTokenizer.from_pretrained(self.name) |
| 131 | if self.tokenizer.chat_template is None: |
| 132 | self.eos += extra_eos_for_direct_completion(dataset) |
| 133 | self.llm = LLM(model=name, max_model_len=2048, **kwargs) |
| 134 | |
| 135 | def is_direct_completion(self) -> bool: |
| 136 | return self.tokenizer.chat_template is None |
| 137 | |
| 138 | def codegen( |
| 139 | self, prompt: str, do_sample: bool = True, num_samples: int = 200 |
| 140 | ) -> List[str]: |
| 141 | if do_sample: |
| 142 | assert self.temperature > 0, "Temperature must be greater than 0!" |
| 143 | batch_size = min(self.batch_size, num_samples) |
| 144 | |
| 145 | vllm_outputs = self.llm.generate( |
| 146 | [prompt] * batch_size, |
| 147 | SamplingParams( |
| 148 | temperature=self.temperature, |
| 149 | max_tokens=self.max_new_tokens, |
| 150 | top_p=0.95 if do_sample else 1.0, |
| 151 | stop=self.eos, |
| 152 | ), |
| 153 | use_tqdm=False, |
| 154 | ) |
| 155 | |
| 156 | gen_strs = [x.outputs[0].text.replace("\t", " ") for x in vllm_outputs] |
| 157 | return gen_strs |
| 158 | |
| 159 | |
| 160 | class GeneralVllmDecoder(VllmDecoder): |
nothing calls this directly
no outgoing calls
no test coverage detected