Views
No views yet
1!pip install -U transformers
2!pip install -U accelerate
3!pip install -U datasets
4!pip install -U peft
5
6!pip install ipywidgets --upgrade
7
8from transformers import (
9 AutoModelForCausalLM,
10 AutoTokenizer,
11 BitsAndBytesConfig,
12)
13from peft import PeftModel
14import torch
15from tqdm import tqdm
16import json
17
18# Hugging Faceで取得したToken(your_token)をこちらに貼る。
19HF_TOKEN = "your_token"
20
21# ベースとなるモデルと学習したLoRAのアダプタ(omnicampusを使用)
22model_id = "models/models--llm-jp--llm-jp-3-13b/snapshots/cd3823f4c1fcbb0ad2e2af46036ab1b0ca13192a"
23adapter_id = "hajim/llm-jp-3-13b-finetune"
24bnb_config = BitsAndBytesConfig(
25 load_in_4bit=True,
26 bnb_4bit_quant_type="nf4",
27 bnb_4bit_compute_dtype=torch.bfloat16,
28)
29model = AutoModelForCausalLM.from_pretrained(
30 model_id,
31 quantization_config=bnb_config,
32 device_map="auto",
33 token = HF_TOKEN
34)
35tokenizer = AutoTokenizer.from_pretrained(model_id, trust_remote_code=True, token = HF_TOKEN)
36model = PeftModel.from_pretrained(model, adapter_id, token = HF_TOKEN)
37
38datasets = []
39with open("./elyza-tasks-100-TV_0.jsonl", "r") as f:
40 item = ""
41 for line in f:
42 line = line.strip()
43 item += line
44 if item.endswith("}"):
45 datasets.append(json.loads(item))
46 item = ""
47results = []
48for data in tqdm(datasets):
49 input = data["input"]
50 prompt = f""" ### 指示
51 {input}
52 ### 回答
53 """
54 tokenized_input = tokenizer.encode(prompt, add_special_tokens=False, return_tensors="pt").to(model.device)
55 attention_mask = torch.ones_like(tokenized_input)
56 with torch.no_grad():
57 outputs = model.generate(
58 tokenized_input,
59 attention_mask=attention_mask,
60 max_new_tokens=100,
61 do_sample=False,
62 repetition_penalty=1.2,
63 pad_token_id=tokenizer.eos_token_id
64 )[0]
65 output = tokenizer.decode(outputs[tokenized_input.size(1):], skip_special_tokens=True)
66 results.append({"task_id": data["task_id"], "input": input, "output": output})
67
68import re
69jsonl_id = re.sub(".*/", "", adapter_id)
70with open(f"./{jsonl_id}-outputs.jsonl", 'w', encoding='utf-8') as f:
71 for result in results:
72 json.dump(result, f, ensure_ascii=False) # ensure_ascii=False for handling non-ASCII characters
73 f.write('\n')