1import torch
2from transformers import AutoTokenizer, AutoModelForCausalLM
3
4device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
5tokenizer = AutoTokenizer.from_pretrained("inu-ai/dolly-japanese-gpt-1b", use_fast=False)
6model = AutoModelForCausalLM.from_pretrained("inu-ai/dolly-japanese-gpt-1b").to(device)
1MAX_ASSISTANT_LENGTH = 100
2MAX_INPUT_LENGTH = 1024
3INPUT_PROMPT = r'<s>\n以下は、タスクを説明する指示と、文脈のある入力の組み合わせです。要求を適切に満たす応答を書きなさい。\n[SEP]\n指示:\n{instruction}\n[SEP]\n入力:\n{input}\n[SEP]\n応答:\n'
4NO_INPUT_PROMPT = r'<s>\n以下は、タスクを説明する指示です。要求を適切に満たす応答を書きなさい。\n[SEP]\n指示:\n{instruction}\n[SEP]\n応答:\n'
5USER_NAME = "User"
6ASSISTANT_NAME = "Assistant"
7
8def prepare_input(role_instruction, conversation_history, new_conversation):
9 instruction = "".join([f"{text} " for text in role_instruction])
10 instruction += " ".join(conversation_history)
11 input_text = f"{USER_NAME}:{new_conversation}"
12
13 return INPUT_PROMPT.format(instruction=instruction, input=input_text)
14
15def format_output(output):
16 output = output.lstrip("<s>").rstrip("</s>").replace("[SEP]", "").replace("\\n", "\n")
17 return output
18
19def generate_response(role_instruction, conversation_history, new_conversation):
20 # 入力トークン数1024におさまるようにする
21 for _ in range(8):
22 input_text = prepare_input(role_instruction, conversation_history, new_conversation)
23 token_ids = tokenizer.encode(input_text, add_special_tokens=False, return_tensors="pt")
24 n = len(token_ids[0])
25 if n + MAX_ASSISTANT_LENGTH <= MAX_INPUT_LENGTH:
26 break
27 else:
28 conversation_history.pop(0)
29 conversation_history.pop(0)
30
31 with torch.no_grad():
32 output_ids = model.generate(
33 token_ids.to(model.device),
34 min_length=n,
35 max_length=min(MAX_INPUT_LENGTH, n + MAX_ASSISTANT_LENGTH),
36 temperature=0.7,
37 repetition_penalty=1.0, # 数値を大きくすると、文字列の繰り返しが減る
38 do_sample=True,
39 pad_token_id=tokenizer.pad_token_id,
40 bos_token_id=tokenizer.bos_token_id,
41 eos_token_id=tokenizer.eos_token_id,
42 bad_words_ids=[[tokenizer.unk_token_id]]
43 )
44
45 output = tokenizer.decode(output_ids.tolist()[0])
46 formatted_output_all = format_output(output)
47
48 response = f"{ASSISTANT_NAME}:{formatted_output_all.split('応答:')[-1].strip()}"
49 conversation_history.append(f"{USER_NAME}:{new_conversation}".replace("\n", "\\n"))
50 conversation_history.append(response.replace("\n", "\\n"))
51
52 return formatted_output_all, response
53
54role_instruction = [
55 f"{USER_NAME}:きみは「ずんだもん」なのだ。東北ずん子の武器である「ずんだアロー」に変身する妖精またはマスコットなのだ。一人称は「ボク」で語尾に「なのだー」を付けてしゃべるのだ。",
56 f"{ASSISTANT_NAME}:了解したのだ。",
57 f"{USER_NAME}:きみは同じ言葉を繰り返さず、何でも正確に要約して答えられるのだ。",
58 f"{ASSISTANT_NAME}:了解したのだ。",
59]
60
61conversation_history = [
62]
63
64questions = [
65 "日本で一番高い山は?",
66 "日本で一番広い湖は?",
67 "冗談を言ってください。",
68 "世界で一番高い山は?",
69 "世界で一番広い湖は?",
70 "最初の質問は何ですか?",
71 "今何問目?",
72 "自己紹介をしてください。",
73]
74
75# 各質問に対して応答を生成して表示
76for question in questions:
77 formatted_output_all, response = generate_response(role_instruction, conversation_history, question)
78 print(f"{USER_NAME}:{question}\n{response}\n---")
User:日本で一番高い山は?
Assistant:富士山
---
User:日本で一番広い湖は?
Assistant:日本列島には、太平洋、日本海、オホーツク海、インド洋、大西洋など、さまざまな海があります。
---
User:冗談を言ってください。
Assistant:冗談を言おう!
---
User:世界で一番高い山は?
Assistant:世界で最も高い山はエベレストです。
---
User:世界で一番広い湖は?
Assistant:太平洋に浮かぶ世界最大の島、ツバル
---
User:最初の質問は何ですか?
Assistant:私が最初の質問をするとするならば、それは「ずんだもん」です。
---
User:今何問目?
Assistant:今何問目ですか?
---
User:自己紹介をしてください。
Assistant:私は「ずんだもん」というあだ名で呼ばれています。
---
1000回の「入力」のような質問を行い、それらに対する「応答」に正解の文字列が含まれるかで評価しています。
一番正答率が高い10エポック目のモデルを選択しました。(やり過ぎたかもしれないです。)
alpacaと同じように、以下のようなフォーマットにしています。
<s>
以下は、タスクを説明する指示と、文脈のある入力の組み合わせです。要求を適切に満たす応答を書きなさい。
[SEP]
指示:
User:きみは「ずんだもん」なのだ。東北ずん子の武器である「ずんだアロー」に変身する妖精またはマスコットなのだ。一人称は「ボク」で語尾に「なのだー」を付けてしゃべるのだ。 Assistant:了解したのだ。 User:きみは同じ言葉を繰り返さず、何でも正確に要約して答えられるのだ。 Assistant:了解したのだ。
[SEP]
入力:
User:日本で一番高い山は?
[SEP]
応答:
富士山
</s>
transformersのコードでtxtファイルを学習する場合、1データ1行のようなので改行コードを一旦
\nに置き換えています。
学習データは
dolly-oasst1-ja.txtです。
また学習データを作った過程のスクリプトとjsonファイルも
train_dataに置いておきます。
※VRAMが足りない場合、optimをadafactorにするとVRAM使用量が減りました。adafactorの場合、learning_rateを1e-03にしてlr_scheduler_typeを削除してと、ChatGPT/GPT-4が言っていました。
venv/Scripts/python.exe transformers/examples/pytorch/language-modeling/run_clm.py ^
--model_name_or_path rinna/japanese-gpt-1b ^
--train_file train_data/dolly-oasst1-ja.txt ^
--output_dir output ^
--do_train ^
--bf16 True ^
--tf32 True ^
--optim adamw_bnb_8bit ^
--num_train_epochs 10 ^
--save_steps 721 ^
--logging_steps 72 ^
--learning_rate 1e-07 ^
--lr_scheduler_type constant ^
--gradient_checkpointing ^
--per_device_train_batch_size 8 ^
--save_safetensors True ^
--logging_dir logs