| 42 | logging.basicConfig(level=logging.INFO) |
| 43 | |
| 44 | class StreamingOpenCodeInterpreter(OpenCodeInterpreter): |
| 45 | streamer: Optional[TextIteratorStreamer] = None |
| 46 | |
| 47 | # overwirte generate function |
| 48 | @torch.inference_mode() |
| 49 | def generate( |
| 50 | self, |
| 51 | prompt: str = "", |
| 52 | max_new_tokens = 1024, |
| 53 | do_sample: bool = False, |
| 54 | top_p: float = 0.95, |
| 55 | top_k: int = 50, |
| 56 | ) -> str: |
| 57 | # Get the model and tokenizer, and tokenize the user text. |
| 58 | |
| 59 | self.streamer = TextIteratorStreamer( |
| 60 | self.tokenizer, skip_prompt=True, Timeout=5 |
| 61 | ) |
| 62 | |
| 63 | inputs = self.tokenizer([prompt], return_tensors="pt", truncation=True, max_length=MAX_INPUT_TOKEN_LENGTH) |
| 64 | inputs = inputs.to(self.model.device) |
| 65 | |
| 66 | kwargs = dict( |
| 67 | **inputs, |
| 68 | streamer=self.streamer, |
| 69 | max_new_tokens=max_new_tokens, |
| 70 | do_sample=do_sample, |
| 71 | top_p=top_p, |
| 72 | top_k=top_k, |
| 73 | eos_token_id=self.tokenizer.eos_token_id |
| 74 | ) |
| 75 | |
| 76 | thread = Thread(target=self.model.generate, kwargs=kwargs) |
| 77 | thread.start() |
| 78 | |
| 79 | return "" |
| 80 | |
| 81 | def save_json(dialog, mode, json_file_path, dialog_id) -> None: |
| 82 | with scheduler.lock: |