Download use.py from VDC-team/VDrontGPT10m-Base: direct link, hf CLI and curl.
- Browser
- Download file 2.57 kB
-
https://huggingface.co/VDC-team/VDrontGPT10m-Base/resolve/main/use.py
- Command line
-
hf download hf://VDC-team/VDrontGPT10m-Base/use.py
-
curl -L -o use.py https://huggingface.co/VDC-team/VDrontGPT10m-Base/resolve/main/use.py
2.57 kB
| import torch | |
| from transformers import AutoModelForCausalLM, AutoTokenizer, TextIteratorStreamer | |
| from threading import Thread | |
| MODEL_PATH = "VDrontGPT10m-Base" | |
| TEMPERATURE = 0.4 | |
| DEVICE = "cuda" if torch.cuda.is_available() else "cpu" | |
| def load_model_and_tokenizer(model_path): | |
| tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=False) | |
| model = AutoModelForCausalLM.from_pretrained( | |
| model_path, | |
| torch_dtype=torch.float16 if DEVICE == "cuda" else torch.float32, | |
| device_map="auto", | |
| trust_remote_code=False | |
| ) | |
| if tokenizer.pad_token is None: | |
| tokenizer.pad_token = tokenizer.eos_token | |
| return model, tokenizer | |
| def generate_stream(model, tokenizer, prompt, temperature=0.4, max_new_tokens=128): | |
| # Text continuation generator without chat history | |
| inputs = tokenizer(prompt, return_tensors="pt", truncation=True, max_length=2048) | |
| inputs = {k: v.to(model.device) for k, v in inputs.items()} | |
| streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True) | |
| generation_kwargs = dict( | |
| **inputs, | |
| max_new_tokens=max_new_tokens, | |
| temperature=temperature, | |
| do_sample=True, | |
| top_p=0.98, | |
| repetition_penalty=1.1, | |
| pad_token_id=tokenizer.pad_token_id, | |
| eos_token_id=tokenizer.eos_token_id, | |
| streamer=streamer, | |
| ) | |
| thread = Thread(target=model.generate, kwargs=generation_kwargs) | |
| thread.start() | |
| for new_text in streamer: | |
| yield new_text | |
| thread.join() | |
| def interactive_chat(model, tokenizer, temperature): | |
| print(f"Simple text continuation (temp={temperature})") | |
| print("Enter the beginning of the text, and the model will continue it.") | |
| print("Commands: 'exit' or 'quit' — exit.") | |
| while True: | |
| try: | |
| user_input = input("\nYou: ").strip() | |
| except (KeyboardInterrupt, EOFError): | |
| print("\nGoodbye!") | |
| break | |
| if user_input.lower() in ["exit", "quit"]: | |
| print("Goodbye!") | |
| break | |
| if not user_input: | |
| continue | |
| print() # empty line before continuation | |
| for token in generate_stream(model, tokenizer, user_input, temperature=temperature): | |
| print(token, end="", flush=True) | |
| print() # finish the line | |
| if __name__ == "__main__": | |
| model, tokenizer = load_model_and_tokenizer(MODEL_PATH) | |
| interactive_chat(model, tokenizer, temperature=TEMPERATURE) |