Views
No views yet
| ANLI-R1 | ANLI-R2 | ANLI-R3 | Avg. |
|---|---|---|---|
| 77.2 | 62.8 | 61.2 | 67.1 |
1from transformers import AutoModelForCausalLM, AutoTokenizer
2from transformers import AutoTokenizer
3from peft import PeftModel
4import torch
5
6base_model_name = 'meta-llama/Llama-3.1-8B-Instruct'
7tokenizer = AutoTokenizer.from_pretrained(base_model_name)
8model = AutoModelForCausalLM.from_pretrained(base_model_name,
9 pad_token_id=tokenizer.eos_token_id,
10 device_map='auto')
11
12lora_model = PeftModel.from_pretrained(model, 'cassuto/Llama-3.1-ANLI-R1-R2-R3-8B-Instruct')
13
14label_str = ['entailment', 'neutral', 'contradiction']
15
16def eval(premise : str, hypothesis : str, device = 'cuda'):
17 input = ("<|start_header_id|>system<|end_header_id|>\n\nBased on the following premise, determine if the hypothesis is entailment, contradiction, or neutral." +
18 "<|eot_id|><|start_header_id|>user<|end_header_id|>\n\n"
19 "<Premise>: " + premise + "\n\n<Hypothesis>: " + hypothesis + "\n\n" +
20 "<|eot_id|><|start_header_id|>assistant<|end_header_id|>")
21 tk = tokenizer(input)
22
23 with torch.no_grad():
24 input_ids = torch.tensor(tk['input_ids']).unsqueeze(0).to(device)
25 out = lora_model.generate(input_ids=input_ids,
26 attention_mask=torch.tensor(tk['attention_mask']).unsqueeze(0).to(device),
27 max_new_tokens=10
28 )
29 print(tokenizer.decode(out[0]))
30 s = tokenizer.decode(out[0][input_ids.shape[-1]:])
31 for lbl, l in enumerate(label_str):
32 if s.find(l) > -1:
33 return lbl
34 else:
35 assert False, 'Invalid model output: ' + s
36
37print(eval("A man is playing a guitar.", "A woman is reading a book."))1from datasets import load_dataset
2import numpy as np
3
4dataset = load_dataset("anli")
5
6model_name = "meta-llama/Llama-3.1-8B-Instruct"
7def out_ckp(r):
8 return f"/path/to/project/Llama-3.1-ANLI-R1-R2-R3-8B-Instruct/checkpoints-r{r}"
9def out_lora_model_fn(r):
10 return f'/path/to/project/Llama-3.1-ANLI-R1-R2-R3-8B-Instruct/lora-r{r}'
11
12from transformers import AutoModelForCausalLM, AutoTokenizer
13
14from transformers import AutoTokenizer, GenerationConfig
15from peft import LoraConfig, get_peft_model, PeftModel
16from trl import SFTConfig, SFTTrainer
17import torch
18from collections.abc import Mapping
19
20label_str = ['entailment', 'neutral', 'contradiction']
21
22def preprocess_function(examples):
23 inputs = ["<|start_header_id|>system<|end_header_id|>\n\nBased on the following premise, determine if the hypothesis is entailment, contradiction, or neutral." +
24 "<|eot_id|><|start_header_id|>user<|end_header_id|>\n\n"
25 "<Premise>: " + p + "\n\n<Hypothesis>: " + h + "\n\n" +
26 "<|eot_id|><|start_header_id|>assistant<|end_header_id|>" + label_str[lbl] + "<|eot_id|>\n" # FIXME remove \n
27 for p, h, lbl in zip(examples["premise"], examples["hypothesis"], examples['label'])]
28
29 model_inputs = {}
30 model_inputs['text'] = inputs
31 return model_inputs
32
33tokenizer = AutoTokenizer.from_pretrained(model_name)
34
35model = AutoModelForCausalLM.from_pretrained(model_name,
36 pad_token_id=tokenizer.eos_token,
37 device_map='auto')
38model.config.use_cache=False
39model.config.pretraining_tp=1
40
41tokenizer.padding_side = "right"
42tokenizer.pad_token = tokenizer.eos_token
43
44for r in range(1,4):
45 print('Round ', r)
46
47 train_data = dataset[f'train_r{r}']
48 val_data = dataset[f'dev_r{r}']
49 train_data = train_data.map(preprocess_function, batched=True,num_proc=8)
50 val_data = val_data.map(preprocess_function, batched=True,num_proc=8)
51
52 training_args = SFTConfig(
53 fp16=True,
54 output_dir=out_ckp(r),
55 learning_rate=1e-4,
56 per_device_train_batch_size=4,
57 per_device_eval_batch_size=1,
58 num_train_epochs=3,
59 logging_steps=10,
60 weight_decay=0,
61 logging_dir=f"./logs-r{r}",
62 save_strategy="epoch",
63 save_total_limit=1,
64 max_seq_length=2048,
65 packing=False,
66 dataset_text_field="text"
67 )
68
69 if r==1:
70 # create LoRA model
71 peft_config = LoraConfig(
72 r=64,
73 lora_alpha=16,
74 lora_dropout=0.1,
75 bias="none",
76 task_type='CAUSAL_LM'
77 )
78 lora_model = get_peft_model(model, peft_config)
79 else:
80 # load the previous trained LoRA part
81 lora_model = PeftModel.from_pretrained(model, out_lora_model_fn(r-1),
82 is_trainable=True)
83
84 trainer = SFTTrainer(
85 model=lora_model,
86 tokenizer=tokenizer,
87 args=training_args,
88 train_dataset=train_data,
89 )
90
91 trainer.train()
92 print(f'saving to "{out_lora_model_fn(r)}"')
93 lora_model.save_pretrained(out_lora_model_fn(r))