| from unsloth import FastLanguageModel |
| import torch |
| import os |
| from transformers import TrainingArguments, TextStreamer |
| from trl import SFTTrainer |
| from datasets import load_dataset |
|
|
| |
| model_name = "unsloth/Meta-Llama-3.1-8B-bnb-4bit" |
| max_seq_length = 2048 |
| load_in_4bit = True |
| dataset_file = "/kaggle/input/datasets/yakoob2345/legal-final/legal_rag_dataset_final.jsonl" |
|
|
| |
| model, tokenizer = FastLanguageModel.from_pretrained( |
| model_name = model_name, |
| max_seq_length = max_seq_length, |
| load_in_4bit = load_in_4bit, |
| ) |
|
|
| |
| model = FastLanguageModel.get_peft_model( |
| model, |
| r = 16, |
| target_modules = ["q_proj", "k_proj", "v_proj", "o_proj", |
| "gate_proj", "up_proj", "down_proj",], |
| lora_alpha = 16, |
| lora_dropout = 0, |
| bias = "none", |
| use_gradient_checkpointing = "unsloth", |
| random_state = 3407, |
| ) |
|
|
| |
| |
| prompt_template = """Below is an instruction that describes a task, paired with an input that provides further context. Write a response that appropriately completes the request. |
| |
| ### Instruction: |
| {} |
| |
| ### Input: |
| {} |
| |
| ### Response: |
| {}""" |
|
|
| EOS_TOKEN = tokenizer.eos_token |
|
|
| def formatting_prompts_func(examples): |
| instructions = examples["instruction"] |
| inputs = examples["input"] |
| outputs = examples["output"] |
| texts = [] |
| for instruction, input, output in zip(instructions, inputs, outputs): |
| |
| text = prompt_template.format(instruction, input, output) + EOS_TOKEN |
| texts.append(text) |
| return { "text" : texts, } |
|
|
| from trl import SFTTrainer, SFTConfig |
|
|
| |
| dataset = load_dataset("json", data_files=dataset_file)["train"] |
| dataset = dataset.train_test_split(test_size=0.1) |
|
|
| train_dataset = dataset["train"].map(formatting_prompts_func, batched=True) |
| eval_dataset = dataset["test"].map(formatting_prompts_func, batched=True) |
|
|
| trainer = SFTTrainer( |
| model=model, |
| tokenizer=tokenizer, |
| train_dataset=train_dataset, |
| eval_dataset=eval_dataset, |
| dataset_text_field="text", |
| max_seq_length=max_seq_length, |
| dataset_num_proc=2, |
| packing=True, |
| args=SFTConfig( |
| per_device_train_batch_size=2, |
| gradient_accumulation_steps=4, |
| warmup_steps=20, |
| max_steps=400, |
| learning_rate=2e-5, |
| fp16=not torch.cuda.is_bf16_supported(), |
| bf16=torch.cuda.is_bf16_supported(), |
| logging_steps=10, |
| eval_strategy="steps", |
| eval_steps=20, |
| save_steps=40, |
| save_total_limit=2, |
| load_best_model_at_end=True, |
| optim="adamw_8bit", |
| weight_decay=0.01, |
| lr_scheduler_type="cosine", |
| seed=3407, |
| output_dir="outputs", |
| report_to="none", |
| logging_first_step=True, |
| gradient_checkpointing_kwargs={"use_reentrant": False}, |
| ), |
| ) |
|
|
| checkpoint_path = "/kaggle/input/datasets/yakoob2345/checkpoint-data/checkpoint-280" |
| |
| trainer.train(resume_from_checkpoint=checkpoint_path) |
|
|
| |
| model.save_pretrained("legal_model_lora") |
| tokenizer.save_pretrained("legal_model_lora") |
| |
|
|