Views
No views yet
1
2import os
3import random
4import numpy as np
5import pandas as pd
6import torch
7from transformers import AutoModelForCausalLM, AutoTokenizer
8from tqdm import tqdm
9import json
10
11# JSONLファイルを読み込む
12file_path = 'elyza-tasks-100-TV_0.jsonl'
13data = pd.read_json(file_path, lines=True)
14
15
16def set_seed(seed):
17 random.seed(seed)
18 os.environ["PYTHONHASHSEED"] = str(seed)
19 np.random.seed(seed)
20 torch.manual_seed(seed)
21 torch.cuda.manual_seed(seed)
22 torch.cuda.manual_seed_all(seed)
23 torch.backends.cudnn.deterministic = True
24 torch.backends.cudnn.benchmark = False
25
26set_seed(42)
27
28
29model_name = "hiroki-rad/google-gemma-2-2b-128-ft-3000"
30
31tokenizer = AutoTokenizer.from_pretrained(model_name)
32
33model = AutoModelForCausalLM.from_pretrained(
34 model_name,
35 torch_dtype="auto",
36 device_map="auto",
37)
38
39def generate_text(data):
40
41 prompt = f"""## 指示:あなたは優秀な日本人の問題解決のエキスパートです。以下のステップで質問に取り組んでください:\n\n1. 質問の種類を特定する(事実確認/推論/創造的回答/計算など)\n2. 重要な情報や制約条件を抽出する\n3. 解決に必要なステップを明確にする\n4. 回答を組み立てる
42 質問をよく読んで、冷静に考え、考えをステップバイステップで考えをまとめてましょう。それをもう一度じっくり考えて、思考のプロセスを整理してください。質問に対して適切な回答を簡潔に出力してください。
43
44
45 質問:{data.input}\n回答:"""
46 # 推論の実行
47 input_ids = tokenizer(prompt, return_tensors="pt").to(model.device)
48 # Remove token_type_ids from the input_ids
49 input_ids.pop('token_type_ids', None)
50 outputs = model.generate(
51 **input_ids,
52 max_new_tokens=2048,
53 do_sample=True,
54 top_p=0.95,
55 temperature=0.9,
56 repetition_penalty=1.1,
57 )
58
59 return tokenizer.decode(outputs[0][len(input_ids['input_ids'][0]):], skip_special_tokens=True)
60
61
62results = []
63for d in tqdm(data.itertuples(), position=0):
64 results.append(generate_text(d))
65
66
67jsonl_data = []
68
69# Iterate through the data and outputs
70for i in range(len(data)):
71 task_id = data.iloc[i]["task_id"] # Access task_id using the index
72 output = results[i]
73
74 # Create a dictionary for each row
75 jsonl_object = {
76 "task_id": task_id,
77 "output": output
78 }
79 jsonl_data.append(jsonl_object)
80
81with open("gemma2-output.jsonl", "w", encoding="utf-8") as outfile:
82 for entry in jsonl_data:
83 # Convert task_id to a regular Python integer before dumping
84 entry["task_id"] = int(entry["task_id"])
85 json.dump(entry, outfile, ensure_ascii=False)
86 outfile.write('\n')
87