Views
No views yet
1# 必要なライブラリをインストール
2!pip install -U bitsandbytes
3!pip install -U transformers
4!pip install -U accelerate
5!pip install -U datasets
6!pip install -U peft
7!pip install ipywidgets --upgrade
8
9# 必要なライブラリを読み込み
10from transformers import (
11 AutoModelForCausalLM,
12 AutoTokenizer,
13 BitsAndBytesConfig,
14)
15from peft import PeftModel
16import torch
17from tqdm import tqdm
18import json
19
20# Hugging Face Token
21from google.colab import userdata
22HF_TOKEN = userdata.get('huggingface_token')
23
24# ベースとなるモデルと学習したLoRAのアダプタ。
25model_id = "llm-jp/llm-jp-3-13b"
26adapter_id = "tash-huggingface/llm-jp-3-13b-finetune"
27
28# QLoRA config
29bnb_config = BitsAndBytesConfig(
30 load_in_4bit=True,
31 bnb_4bit_quant_type="nf4",
32 bnb_4bit_compute_dtype=torch.bfloat16,
33)
34
35# Load model
36model = AutoModelForCausalLM.from_pretrained(
37 model_id,
38 quantization_config=bnb_config,
39 device_map="auto",
40 token = HF_TOKEN
41)
42
43# Load tokenizer
44tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True, token = HF_TOKEN)
45
46# 元のモデルにLoRAのアダプタを統合。
47model = PeftModel.from_pretrained(model, adapter_id, token = HF_TOKEN)
48
49# データセットの読み込み。
50datasets = []
51with open("./elyza-tasks-100-TV_0.jsonl", "r") as f:
52 item = ""
53 for line in f:
54 line = line.strip()
55 item += line
56 if item.endswith("}"):
57 datasets.append(json.loads(item))
58 item = ""
59
60results = []
61for data in tqdm(datasets):
62 input = data["input"]
63 prompt = f"""### 指示
64 {input}
65 ### 回答
66 """
67 tokenized_input = tokenizer.encode(prompt, add_special_tokens=False, return_tensors="pt").to(model.device)
68 attention_mask = torch.ones_like(tokenized_input)
69 with torch.no_grad():
70 outputs = model.generate(
71 tokenized_input,
72 attention_mask=attention_mask,
73 max_new_tokens=100,
74 do_sample=False,
75 repetition_penalty=1.2,
76 pad_token_id=tokenizer.eos_token_id
77 )[0]
78 output = tokenizer.decode(outputs[tokenized_input.size(1):], skip_special_tokens=True)
79 results.append({"task_id": data["task_id"], "input": input, "output": output})
80
81# 結果をjsonlで出力
82import re
83jsonl_id = re.sub(".*/", "", adapter_id)
84with open(f"./{jsonl_id}-outputs.jsonl", 'w', encoding='utf-8') as f:
85 for result in results:
86 json.dump(result, f, ensure_ascii=False) # ensure_ascii=False for handling non-ASCII characters
87 f.write('\n')