import unsloth import torch import datasets from trl import SFTTrainer from unsloth import FastLanguageModel from transformers import TrainingArguments max_seq_length = 512 # Can increase for longer reasoning traces lora_rank = 16 # Larger rank = smarter, but slower orig_model_path = 'OpenLLM-Ro/RoLlama3.1-8b-Instruct' mask = 0x3F NUM_EPOCHS=3 out_model_name = f'roLl31I-200K-{mask:04X}-EP{NUM_EPOCHS}-1per' COUNT = 200000 #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/gram-errors-3F-ph_sent') ds_train_orig = dsd['phrases'] #.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['text'].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'[INST]\n{user_message_1} [/INST] {model_reply_1}\n' 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") 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