File size: 4,904 Bytes
14116b4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6c85d58
14116b4
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
import unsloth
import torch
import datasets
from trl import SFTTrainer
from unsloth import FastLanguageModel
from transformers import TrainingArguments

max_seq_length = 1024 # Can increase for longer reasoning traces
lora_rank = 16 # Larger rank = smarter, but slower
orig_model_path = 'OpenLLM-Ro/RoLlama3.1-8b-Instruct'
mask = 0x3
NUM_EPOCHS=2
out_model_name = f'roLl31I-RRT_PH-{mask:04X}-EP{NUM_EPOCHS}'


COUNT = None

#dsd = datasets.load_dataset('hartular/rrt-grammatical_errors-split')
#ds_train_orig = dsd['train'].filter(lambda ex: (0x01 << ex['error_class']) & mask)
#ds_train_orig.rename_column('input', 'text')

dsd = datasets.load_dataset('hartular/rrt-agree-phraseonly-train')
ds_train_orig = dsd['train'] #.filter(lambda ex: (0x01 << ex['error_class']) & mask)


# ds_orig = datasets.load_dataset('hartular/gram-err-36DB-train-2per')

# transform to good_good and good_bad pairs
# ds_dict = ds_orig['train']

# for split in ds_orig.keys():
#     orig_data = ds_orig[split].to_list()

data_list = []
for d in ds_train_orig.to_list():
#    data_list.extend([{'input':d['good_text' if is_good else 'bad_text'],
#                       'response':d['good_text']} for is_good in (False, True)])
    data_list.append({'input':d['input'].replace('\xad', ''), 'response':d['response']})

#                        {'input':d['bad_text'], 'response':'0'}])

#data_list.sort(key=lambda d: len(d['input']))

ds_train = datasets.Dataset.from_list(data_list)

model, tokenizer = FastLanguageModel.from_pretrained(
    model_name = orig_model_path,
    max_seq_length = max_seq_length,
    load_in_4bit = True, # False for LoRA 16bit
    fast_inference = True, # Enable vLLM fast inference
    max_lora_rank = lora_rank,
    gpu_memory_utilization = 0.6, # Reduce if out of memory
)

model = FastLanguageModel.get_peft_model(
    model,
    r = lora_rank, # Choose any number > 0 ! Suggested 8, 16, 32, 64, 128
    target_modules = [
        "q_proj", "k_proj", "v_proj", "o_proj",
        "gate_proj", "up_proj", "down_proj",
    ], # Remove QKVO if out of memory
    lora_alpha = lora_rank,
    use_gradient_checkpointing = "unsloth", # Enable long context finetuning
    random_state = 1,
)

import json

def preprocess_function(ex) -> list[str]:
    return [tokenizer.apply_chat_template(
        conversation=[
            # {'role':'system', 'content':'Ești un automat care răspunde cu 1 dacă enunțul pe care l-a primit este corect gramatical și răspunde cu 0 dacă enunțul pe care l-a primit nu este corect gramatical.'},
            {"role":"user", "content":in_str},
            {"role":"assistant", "content": res_str},
        ], tokenize=False, max_seq_length=max_seq_length, truncate=True) for in_str, res_str in zip(ex['input'], ex['response'])]

#ds_train = ds_train.shuffle()
if COUNT:
    ds_train = ds_train.select(range(COUNT))

def preprocess_function_llama2(ex) -> list[str]:
    return [
        f'<s>[INST]\n{user_message_1} [/INST] {model_reply_1}\n</s>' for user_message_1, model_reply_1 in
            zip(ex['text'], ex['response'])
    ]
args=TrainingArguments(
        learning_rate=3e-4,
        lr_scheduler_type="linear",
        per_device_train_batch_size=8,
        gradient_accumulation_steps=2,
        num_train_epochs=NUM_EPOCHS,
        fp16=not unsloth.is_bfloat16_supported(),
        bf16=unsloth.is_bfloat16_supported(),
        logging_steps=1,
        optim="adamw_8bit",
        weight_decay=0.01,
        warmup_steps=10,
        output_dir=out_model_name,
        seed=0,
)

trainer=SFTTrainer(model=model,
                   tokenizer=tokenizer,
                   formatting_func=preprocess_function,
                   train_dataset=ds_train,
                   #dataset_text_field="text",
                   max_seq_length=max_seq_length,
                   args=args,
    # dataset_num_proc=2,
    # packing=True,
)

trainer.train()
# model.save_pretrained_merged(out_model_name, tokenizer, save_method="lora")
model.save_pretrained_merged(out_model_name, tokenizer, save_method="merged_16bit")
# model.push_to_hub_merged("hartular/" + out_model_name, tokenizer, save_method = "merged_16bit",)
def get_response(msg : str, with_system = False, **kwargs) -> str:
    # to_dev = kwargs.get('to_dev')
    msg = [{'role':'user', 'content':msg}]
    if with_system:
        msg = [{'role':'system', 'content':'Ești un automat care răspunde cu 1 dacă enunțul pe care l-a primit este corect gramatical și răspunde cu 0 dacă enunțul pe care l-a primit nu este corect gramatical.'},] + msg
    inputs = tokenizer.apply_chat_template(msg, tokenize=True, return_tensors="pt",).to('cuda:0')
    out = model.generate(input_ids=inputs, max_new_tokens=128, use_cache=True)
    out_str = tokenizer.decode(out[0])
    try:
        out_str = out_str.split('<|end_header_id|>')[-1].strip('<|eot_id|>').strip()
    except:
        pass
    return out_str