Views
No views yet
1# ライブラリのインストール
2!pip install -U langchain-community langchain-huggingface vllm triton fugashi unidic-lite
3# インストール
4import json
5import pandas as pd
6from tqdm import tqdm
7from transformers import pipeline
8from langchain_community.llms import VLLM
9from langchain.prompts import PromptTemplate
10from langchain_core.runnables import RunnablePassthrough
11from langchain.schema.output_parser import StrOutputParser
12
13
14# GitHub repositoryのclone
15!git clone https://github.com/y-hiroki-radiotech/llm-final-task.git
16%cd llm-final-task
17
18# タスク別に設定したプロンプトを使うために、PromptStockクラスをインスタンス化
19from prompt import PromptStock
20prompt_stock = PromptStock()
21
22# データのpandas形式で準備する
23file_path = 'elyza-tasks-100-TV_0.jsonl' # ここにjsonlを指定する
24data = pd.read_json(file_path, lines=True)
25
26# データのinputに対して、タスクラベルを与える。タスクを8分類してある。
27model_name = "hiroki-rad/bert-base-classification-ft"
28classify_pipe = pipeline(model=model_name, device="cuda:0")
29
30results: list[dict[str, float | str]] = []
31for example in data.itertuples():
32 # モデルの予測結果を取得
33 model_prediction = classify_pipe(example.input)[0]
34 # 正解のラベルIDをラベル名に変換
35 results.append( model_prediction["label"])
36
37data["label"] = results
38
39# タスク回答のためのモデルをvLLMを使ってインストール
40model_name = "hiroki-rad/llm-jp-llm-jp-3-13b-16-ft"
41
42llm = VLLM(model=model_name,
43 quantization="awq")
44
45# テンプレートの作成
46template = """
47ユーザー: 質問を良く読んで、適切な回答をしてください。
48{context}
49質問:{input}
50回答:"""
51
52prompt = PromptTemplate(
53 template=template,
54 input_variables=["context", "input"],
55 template_format="f-string"
56)
57# chainの作成
58vllm_chain = prompt | llm
59
60chain = (
61 RunnablePassthrough()
62 | vllm_chain
63 | StrOutputParser()
64)
65
66outputs = []
67total_rows = len(data)
68with tqdm(total=total_rows,
69 desc="Processing rows",
70 position=0,
71 leave=True
72 ) as pbar:
73 for row in data.itertuples():
74 prompt_string = prompt_stock.get_prompt(row.label)
75
76 input_dict = {
77 "context": prompt_string,
78 "input": row.input
79 }
80
81 output = chain.invoke(input_dict)
82 outputs.append(output)
83
84 pbar.update(1)
85
86# 出力
87jsonl_data = []
88
89for i in range(len(data)):
90 task_id = data.iloc[i]["task_id"]
91 output = outputs[i]
92
93 jsonl_object = {
94 "task_id": task_id,
95 "output": output
96 }
97 jsonl_data.append(jsonl_object)
98
99with open("output.jsonl", "w") as outfile:
100 for entry in jsonl_data:
101 entry["task_id"] = int(entry["task_id"])
102 json.dump(entry, outfile)
103 outfile.write('\n')