Views
No views yet
1!pip install -U bitsandbytes
2!pip install -U transformers
3!pip install -U accelerate
4!pip install -U datasets
5!pip install -U peft
6
7!pip install ipywidgets --upgrade
8
9from transformers import (
10 AutoModelForCausalLM,
11 AutoTokenizer,
12 BitsAndBytesConfig,
13)
14from peft import PeftModel
15import torch
16from tqdm import tqdm
17import json
18
19HF_TOKEN = "Hugging Face Token"
20
21model_id = "models/models--llm-jp--llm-jp-3-13b/snapshots/cd3823f4c1fcbb0ad2e2af46036ab1b0ca13192a"
22adapter_id = "KEN-1Q80/llm-jp-3-13b-finetune02"
23
24bnb_config = BitsAndBytesConfig(
25 load_in_4bit=True,
26 bnb_4bit_quant_type="nf4",
27 bnb_4bit_compute_dtype=torch.bfloat16,
28)
29
30model = AutoModelForCausalLM.from_pretrained(
31 model_id,
32 quantization_config=bnb_config,
33 device_map="auto",
34 token = HF_TOKEN
35)
36
37tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True, token = HF_TOKEN)
38
39model = PeftModel.from_pretrained(model, adapter_id, token = HF_TOKEN)
40
41datasets = []
42with open("./elyza-tasks-100-TV_0.jsonl", "r") as f:
43 item = ""
44 for line in f:
45 line = line.strip()
46 item += line
47 if item.endswith("}"):
48 datasets.append(json.loads(item))
49 item = ""
50
51results = []
52for data in tqdm(datasets):
53
54 input = data["input"]
55
56 prompt = f"""### 指示
57 {input}
58 ### 回答
59 """
60
61 tokenized_input = tokenizer.encode(prompt, add_special_tokens=False, return_tensors="pt").to(model.device)
62 attention_mask = torch.ones_like(tokenized_input)
63 with torch.no_grad():
64 outputs = model.generate(
65 tokenized_input,
66 attention_mask=attention_mask,
67 max_new_tokens=100,
68 do_sample=False,
69 repetition_penalty=1.2,
70 pad_token_id=tokenizer.eos_token_id
71 )[0]
72 output = tokenizer.decode(outputs[tokenized_input.size(1):], skip_special_tokens=True)
73
74 results.append({"task_id": data["task_id"], "input": input, "output": output})
75
76import re
77jsonl_id = re.sub(".*/", "", adapter_id)
78with open(f"./{jsonl_id}-outputs.jsonl", 'w', encoding='utf-8') as f:
79 for result in results:
80 json.dump(result, f, ensure_ascii=False)
81 f.write('\n')