Views
No views yet
1%pip install vllm
2%pip install accelerate
3%pip install flash-attn
4# %pip install packaging==24.1 ### packaging==24.1にしないとエラーになる!! ###
5%pip install huggingface_hub
6%pip install "bitsandbytes>=0.44.0"1import json
2import os
3import vllm ### packaging==24.1にしないとエラーになる!! ###
4from huggingface_hub import snapshot_download
5from transformers import (
6 AutoModelForCausalLM,
7 AutoTokenizer,
8)
9from vllm.lora.request import LoRARequest
10
11HF_TOKEN = os.getenv("HF_TOKEN")
12if HF_TOKEN is None:
13 raise EnvironmentError("環境変数 'HF_TOKEN' が設定されていません。")
14
15MODEL_NAME = "llm-jp/llm-jp-3-13b"
16adapter_id = "YY1128/llm-jp-3-13b-ft1"
17
18lora_path = snapshot_download(repo_id=adapter_id, use_auth_token=HF_TOKEN)1llm = vllm.LLM(
2 MODEL_NAME,
3 gpu_memory_utilization=0.95,
4 trust_remote_code=True,
5 enable_lora=True,
6 max_model_len=2048,
7 enforce_eager=True,
8 quantization="bitsandbytes", # 4bit量子化を有効にする
9 load_format="bitsandbytes", # 4bit量子化を有効にする
10 max_lora_rank=64
11)
12tokenizer = llm.get_tokenizer()
13bos = tokenizer.bos_token1def read_jsonl(file_path):
2 """
3 JSONLファイルを読み込み、辞書のリストを返す関数
4
5 Args:
6 file_path (str): JSONLファイルのパス
7
8 Returns:
9 list: 辞書のリスト
10 """
11 data = []
12 with open(file_path, 'r', encoding='utf-8') as file:
13 for line in file:
14 # 各行をJSON形式でパースしてリストに追加
15 data.append(json.loads(line.strip()))
16 return data
17
18tasks = read_jsonl("elyza-tasks-100-TV_0.jsonl")1inputs = [d["input"] for d in tasks]
2
3prompt_token_ids = [
4 bos + "## 指示\n" + input + "\n" + "## 回答\n"
5 for input in inputs
6]
7print(prompt_token_ids[0])1<s>## 指示
2野球選手が今シーズン活躍するために取り組むべき5つのことを教えてください。
3## 回答1sampling_params = vllm.SamplingParams(
2 best_of=8,
3 max_tokens=700,
4 top_p=0.9,
5 min_p=0.08,
6 seed=42,
7 temperature=1.0,
8 repetition_penalty=1.1,
9) best,max_tokens=700, top_p=0.9, min_p=0.08, seed=42, temperature=1.0,repetition_penalty=1.1
10)1outputs = llm.generate(
2 prompts=prompt_token_ids,
3 sampling_params=sampling_params,
4 lora_request=LoRARequest("lora_1", 1, lora_path), # 第一引数は任意の文字列でよい
5)
6
7import json
8
9data = [
10 {
11 "task_id": i,
12 "input": inputs[i], # あってもなくてもよい
13 "output": outputs[i].outputs[0].text.strip(),
14 }
15 for i in range(len(tasks))
16]
17file_path = "output.jsonl"
18with open(file_path, "w", encoding="utf-8") as file:
19 for entry in data:
20 json.dump(entry, file, ensure_ascii=False)
21 file.write("\n")