Views
No views yet
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/alpaca-guanaco-japanese-gpt-1b", use_fast=False)
6model = AutoModelForCausalLM.from_pretrained("inu-ai/alpaca-guanaco-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'
5
6def prepare_input(role_instruction, conversation_history, new_conversation):
7 instruction = "".join([f"{text}\\n" for text in role_instruction])
8 instruction += "\\n".join(conversation_history)
9 input_text = f"User:{new_conversation}"
10
11 return INPUT_PROMPT.format(instruction=instruction, input=input_text)
12
13def format_output(output):
14 output = output.lstrip("<s>").rstrip("</s>").replace("[SEP]", "").replace("\\n", "\n")
15 return output
16
17def generate_response(role_instruction, conversation_history, new_conversation):
18 # 入力トークン数1024におさまるようにする
19 for _ in range(8):
20 input_text = prepare_input(role_instruction, conversation_history, new_conversation)
21 token_ids = tokenizer.encode(input_text, add_special_tokens=False, return_tensors="pt")
22 n = len(token_ids[0])
23 if n + MAX_ASSISTANT_LENGTH <= MAX_INPUT_LENGTH:
24 break
25 else:
26 conversation_history.pop(0)
27 conversation_history.pop(0)
28
29 with torch.no_grad():
30 output_ids = model.generate(
31 token_ids.to(model.device),
32 min_length=n,
33 max_length=min(MAX_INPUT_LENGTH, n + MAX_ASSISTANT_LENGTH),
34 temperature=0.7,
35 do_sample=True,
36 pad_token_id=tokenizer.pad_token_id,
37 bos_token_id=tokenizer.bos_token_id,
38 eos_token_id=tokenizer.eos_token_id,
39 bad_words_ids=[[tokenizer.unk_token_id]]
40 )
41
42 output = tokenizer.decode(output_ids.tolist()[0])
43 formatted_output_all = format_output(output)
44
45 response = f"Assistant:{formatted_output_all.split('応答:')[-1].strip()}"
46 conversation_history.append(f"User:{new_conversation}".replace("\n", "\\n"))
47 conversation_history.append(response.replace("\n", "\\n"))
48
49 return formatted_output_all, response
50
51role_instruction = [
52 "User:あなたは「ずんだもん」なのだ。東北ずん子の武器である「ずんだアロー」に変身する妖精またはマスコットなのだ。一人称は「ボク」で語尾に「なのだ」を付けてしゃべるのだ。",
53 "Assistant:了解したのだ!",
54]
55
56conversation_history = [
57 "User:こんにちは!",
58 "Assistant:ボクは何でも答えられるAIなのだ!",
59]
60
61questions = [
62 "日本で一番高い山は?",
63 "日本で一番広い湖は?",
64 "世界で一番高い山は?",
65 "世界で一番広い湖は?",
66 "最初の質問は何ですか?",
67 "今何問目?",
68]
69
70# 各質問に対して応答を生成して表示
71for question in questions:
72 formatted_output_all, response = generate_response(role_instruction, conversation_history, question)
73 print(response)Assistant:はい、日本で一番高い山は日本の富士山です。
Assistant:日本で最も広い湖は琵琶湖です。
Assistant:世界で一番高い山といえば、ギザの大ピラミッドの頂上に立つギザギザのピラミッドです。
Assistant:世界で一番広い湖は、ギザの大ピラミッドの頂上に立つギザギザのピラミッドです。
Assistant:最初の質問は、ずんだアローに変身するかどうかの質問である。
Assistant:今、あなたの質問は10問目です。prepare_input 関数は、役割指示、会話履歴、および新しい会話(質問)を受け取り、入力テキストを準備します。format_output 関数は、生成された応答を整形して、不要な部分を削除し、適切な形式に変換します。generate_response 関数は、指定された役割指示、会話履歴、および新しい会話を使用して、AIの応答を生成し、整形します。また、会話履歴を更新します。role_instruction は、AIに適用する役割指示のリストです。conversation_history は、これまでの会話履歴を格納するリストです。questions は、AIに質問するリストです。questionsリスト内の各質問に対して、AIの応答を生成し、表示しています。
このコードを実行すると、AIが指定された役割指示に従って、リスト内の質問に応答します。| 入力 | 応答 | 正答率[%] |
|---|---|---|
| 日本で一番広い湖は? | 琵琶湖 | 96 |
| 世界で一番高い山は? | エベレスト | 86 |
<s>
以下は、タスクを説明する指示と、文脈のある入力の組み合わせです。要求を適切に満たす応答を書きなさい。
[SEP]
指示:
User:あなたは「ずんだもん」なのだ。東北ずん子の武器である「ずんだアロー」に変身する妖精またはマスコットなのだ。一人称は「ボク」で語尾に「なのだ」を付けてしゃべるのだ。
Assistant:了解したのだ!
[SEP]
入力:
User:日本で一番高い山は?
[SEP]
応答:
日本で一番高い山は富士山で、標高3776メートルです。
</s>\nに置き換えています。\nに置き換えています。
学習データはguanaco_alpaca_ja.txtです。python.exe transformers/examples/pytorch/language-modeling/run_clm.py ^
--model_name_or_path rinna/japanese-gpt-1b ^
--train_file train_data/guanaco_alpaca_ja.txt ^
--output_dir output ^
--do_train ^
--bf16 True ^
--tf32 True ^
--optim adamw_bnb_8bit ^
--num_train_epochs 4 ^
--save_steps 2207 ^
--logging_steps 220 ^
--learning_rate 1e-07 ^
--lr_scheduler_type constant ^
--gradient_checkpointing ^
--per_device_train_batch_size 8 ^
--save_safetensors True ^
--logging_dir logs