1from transformers import AutoTokenizer, AutoModelForCausalLM
2import torch
3
4# モデルとトークナイザーの読み込み
5model = AutoModelForCausalLM.from_pretrained(
6 "eyepyon/rc6_elyza_quen32b_fine_merged_v1",
7 torch_dtype=torch.bfloat16, # H100最適化
8 device_map="auto",
9 trust_remote_code=True
10)
11
12tokenizer = AutoTokenizer.from_pretrained(
13 "eyepyon/rc6_elyza_quen32b_fine_merged_v1",
14 trust_remote_code=True
15)
16
17# プロンプトテンプレート
18def create_prompt(question, choices=None):
19 if choices:
20 prompt = f'''<|im_start|>user
21以下の司法試験問題について、思考プロセスを示しながら回答してください。
22
23問題:
24{question}
25
26選択肢:
27{choices}
28
29段階的に分析して正解を導いてください。
30<|im_end|>
31<|im_start|>assistant'''
32 else:
33 prompt = f'''<|im_start|>user
34以下の司法試験問題について、思考プロセスを示しながら回答してください。
35
36問題:
37{question}
38
39法的根拠を示しながら論述してください。
40<|im_start|>
41<|im_start|>assistant'''
42
43 return prompt.format(question=question, choices=choices)
44
45# 推論実行
46question = "憲法第21条の表現の自由について説明してください。"
47prompt = create_prompt(question)
48
49inputs = tokenizer(prompt, return_tensors="pt")
50with torch.no_grad():
51 outputs = model.generate(
52 **inputs,
53 max_new_tokens=512,
54 temperature=0.3,
55 do_sample=True,
56 top_p=0.9,
57 pad_token_id=tokenizer.eos_token_id,
58 eos_token_id=tokenizer.eos_token_id
59 )
60
61response = tokenizer.decode(outputs[0], skip_special_tokens=True)
62print(response)
1# H100での最適化設定
2import torch.backends.cuda
3
4# TensorFloat-32有効化
5torch.backends.cuda.matmul.allow_tf32 = True
6torch.backends.cudnn.allow_tf32 = True
7
8# コンパイル最適化(PyTorch 2.0以上)
9model = torch.compile(model, mode="max-autotune")
10
11# バッチ推論
12questions = [
13 "民法第176条について説明してください。",
14 "刑法における故意の概念について説明してください。"
15]
16
17prompts = [create_prompt(q) for q in questions]
18inputs = tokenizer(prompts, return_tensors="pt", padding=True)
19
20with torch.no_grad():
21 outputs = model.generate(
22 **inputs,
23 max_new_tokens=512,
24 temperature=0.3,
25 do_sample=True,
26 num_beams=1, # 高速化
27 pad_token_id=tokenizer.eos_token_id
28 )
1# 必須ライブラリ
2pip install torch>=2.0.0
3pip install transformers>=4.36.0
4pip install accelerate>=0.25.0
5pip install bitsandbytes>=0.41.0
6
7# 高速化用(オプション)
8pip install flash-attn --no-build-isolation