Views
No views yet
1
2
3from peft import PeftModel, PeftConfig
4from transformers import AutoModelForCausalLM
5from transformers import AutoTokenizer
6import torch
7
8device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
9
10config = PeftConfig.from_pretrained("Ashishkr/llama2-qrecc-context-resolution")
11model = AutoModelForCausalLM.from_pretrained("meta-llama/Llama-2-7b-hf")
12model = PeftModel.from_pretrained(model, "Ashishkr/llama2-qrecc-context-resolution").to(device)
13tokenizer = AutoTokenizer.from_pretrained("meta-llama/Llama-2-7b-hf")
14
15def response_generate(
16 model: AutoModelForCausalLM,
17 tokenizer: AutoTokenizer,
18 prompt: str,
19 max_new_tokens: int = 128,
20 temperature: float = 0.7,
21):
22 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
23
24 inputs = tokenizer(
25 [prompt],
26 return_tensors="pt",
27 return_token_type_ids=False,
28 ).to(
29 device
30 )
31
32 with torch.autocast("cuda", dtype=torch.bfloat16):
33 response = model.generate(
34 **inputs,
35 max_new_tokens=max_new_tokens,
36 temperature=temperature,
37 return_dict_in_generate=True,
38 eos_token_id=tokenizer.eos_token_id,
39 pad_token_id=tokenizer.pad_token_id,
40 )
41
42 decoded_output = tokenizer.decode(
43 response["sequences"][0],
44 skip_special_tokens=True,
45 )
46
47 return decoded_output
48
49prompt = """ Strictly use the context provided, to generate the repsonse. No additional information to be added. Re-write the user query using the context .
50>>CONTEXT<<Where did jessica go to school? Where did she work at?>>USER<<What did she do next for work?>>REWRITE<<"""
51
52response = response_generate(
53 model,
54 tokenizer,
55 prompt,
56 max_new_tokens=20,
57 temperature=0.1,
58)
59
60print(response)
61