From 58f3a6e29d924e0ed54a3fdf2ffeb3492ebda0cd Mon Sep 17 00:00:00 2001 From: Roland Tannous <115670425+rolandtannous@users.noreply.github.com> Date: Tue, 3 Jun 2025 06:32:57 +0300 Subject: [PATCH] reroute merge logic language models + comprehensive tests + eval kits (#2673) --- .../test_merge_model_perplexity_llama-3.2.py | 218 +++++ .../test_merge_model_perplexity_mistral.py | 298 +++++++ .../test_merge_model_perplexity_phi_4.py | 222 +++++ ...st_merged_model_perplexity_llama-3.1-8b.py | 224 +++++ .../test_merged_model_perplexity_qwen_2.5.py | 287 +++++++ .../test_push_to_hub_merged.py | 178 ++++ ...t_push_to_hub_merged_sharded_index_file.py | 197 +++++ .../test_save_merged_grpo_model.py | 794 ++++++++++++++++++ .../test_index_file_sharded_model.py | 289 +++++++ .../vision_models/test_push_to_hub_merged.py | 268 ++++++ ...t_save_merge_vision_model_ocr_benchmark.py | 254 ++++++ tests/utils/aime_eval.md | 264 ++++++ tests/utils/aime_eval.py | 495 +++++++++++ tests/utils/cleanup_utils.py | 209 +++++ tests/utils/ocr_eval.md | 109 +++ tests/utils/ocr_eval.py | 352 ++++++++ tests/utils/perplexity_eval.md | 20 + tests/utils/perplexity_eval.py | 75 ++ unsloth/save.py | 14 +- 19 files changed, 4760 insertions(+), 7 deletions(-) create mode 100644 tests/saving/language_models/test_merge_model_perplexity_llama-3.2.py create mode 100644 tests/saving/language_models/test_merge_model_perplexity_mistral.py create mode 100644 tests/saving/language_models/test_merge_model_perplexity_phi_4.py create mode 100644 tests/saving/language_models/test_merged_model_perplexity_llama-3.1-8b.py create mode 100644 tests/saving/language_models/test_merged_model_perplexity_qwen_2.5.py create mode 100644 tests/saving/language_models/test_push_to_hub_merged.py create mode 100644 tests/saving/language_models/test_push_to_hub_merged_sharded_index_file.py create mode 100644 tests/saving/language_models/test_save_merged_grpo_model.py create mode 100644 tests/saving/vision_models/test_index_file_sharded_model.py create mode 100644 tests/saving/vision_models/test_push_to_hub_merged.py create mode 100644 tests/saving/vision_models/test_save_merge_vision_model_ocr_benchmark.py create mode 100644 tests/utils/aime_eval.md create mode 100644 tests/utils/aime_eval.py create mode 100644 tests/utils/cleanup_utils.py create mode 100644 tests/utils/ocr_eval.md create mode 100644 tests/utils/ocr_eval.py create mode 100644 tests/utils/perplexity_eval.md create mode 100644 tests/utils/perplexity_eval.py diff --git a/tests/saving/language_models/test_merge_model_perplexity_llama-3.2.py b/tests/saving/language_models/test_merge_model_perplexity_llama-3.2.py new file mode 100644 index 0000000000..384da30e85 --- /dev/null +++ b/tests/saving/language_models/test_merge_model_perplexity_llama-3.2.py @@ -0,0 +1,218 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc + +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + +# Define helper functions outside of main +def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + +def load_and_compute_8bit_ppl(result_queue, load_in_4bit=False, load_in_8bit=False): + """Load model and compute perplexity in subprocess""" + from unsloth import FastLanguageModel + from unsloth.chat_templates import get_chat_template + from tests.utils.perplexity_eval import ppl_model + + # Load model + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_llama_text_model", + max_seq_length=2048, + load_in_4bit=load_in_4bit, + load_in_8bit=load_in_8bit, + ) + # Set up tokenizer + merged_tokenizer = get_chat_template( + merged_tokenizer, + chat_template="llama-3.1", + ) + + # Load dataset fresh in subprocess + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + # Format the dataset + def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [merged_tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + # Compute perplexity using the passed dataset + ppl_value = ppl_model(merged_model, merged_tokenizer, dataset_ppl) + + + # IMPORTANT: Convert to Python float if it's a tensor + if torch.is_tensor(ppl_value): + ppl_value = ppl_value.cpu().item() # Move to CPU and convert to Python scalar + elif hasattr(ppl_value, 'item'): + ppl_value = ppl_value.item() # Convert numpy or other array types + else: + ppl_value = float(ppl_value) # Ensure it's a float + + # Return only the perplexity value + result_queue.put(ppl_value) + + # Clean up + del merged_model + del merged_tokenizer + del dataset_ppl + torch.cuda.empty_cache() + gc.collect() + +# Main execution code should be wrapped in this guard +if __name__ == "__main__": + mp.set_start_method('spawn', force=True) + + if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' + else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + + model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/Llama-3.2-3B-Instruct", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation + ) + + tokenizer = get_chat_template( + tokenizer, + chat_template="llama-3.1", + ) + + from unsloth.chat_templates import standardize_sharegpt + dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + dataset_train = dataset_train.map(formatting_prompts_func, batched=True) + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + + model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, + ) + + from unsloth import is_bfloat16_supported + + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer), + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=10, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), + ) + + from unsloth.chat_templates import train_on_responses_only + trainer = train_on_responses_only( + trainer, + instruction_part="<|start_header_id|>user<|end_header_id|>\n\n", + response_part="<|start_header_id|>assistant<|end_header_id|>\n\n", + ) + + # run training + trainer_stats = trainer.train() + + add_to_comparison("Qlora model", ppl_model(model, tokenizer, dataset_ppl)) + + # saving and merging the model to local disk + print("merge and save to local disk") + model.save_pretrained_merged( + save_directory='./unsloth_out/merged_llama_text_model', + tokenizer=tokenizer + ) + + # print("cleaning") + # del model + # del tokenizer + # torch.cuda.empty_cache() + # gc.collect() + + # load model from local disk and test + print("Loading merged model in 4 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_llama_text_model", + max_seq_length=2048, + load_in_4bit=True, + load_in_8bit=False, + ) + + add_to_comparison("merged model load 4bit", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + + print("Computing 8-bit model perplexity in subprocess...") + result_queue = mp.Queue() + p = mp.Process(target=load_and_compute_8bit_ppl, args=(result_queue, False, True)) + p.start() + p.join() + + ppl_8bit = result_queue.get() + add_to_comparison("merged model loaded 8bits", ppl_8bit) + + print("Loading merged model in 16 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_llama_text_model", + max_seq_length=2048, + load_in_4bit=False, + load_in_8bit=False, + ) + + add_to_comparison("merged model loaded 16bits", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + print_model_comparison() + + # final cleanup + safe_remove_directory("./outputs") + safe_remove_directory("./unsloth_compiled_cache") + safe_remove_directory("./unsloth_out") diff --git a/tests/saving/language_models/test_merge_model_perplexity_mistral.py b/tests/saving/language_models/test_merge_model_perplexity_mistral.py new file mode 100644 index 0000000000..f08a3c1fe0 --- /dev/null +++ b/tests/saving/language_models/test_merge_model_perplexity_mistral.py @@ -0,0 +1,298 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc + +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + + + + + + +def load_and_compute_8bit_ppl(result_queue, load_in_4bit=False, load_in_8bit=False): + """Load model and compute perplexity in subprocess""" + from unsloth import FastLanguageModel + from tests.utils.perplexity_eval import ppl_model + + # Load model + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_mistral_text_model", + max_seq_length=2048, + load_in_4bit=load_in_4bit, + load_in_8bit=load_in_8bit, + ) + # Set up tokenizer + # merged_tokenizer = get_chat_template( + # merged_tokenizer, + # chat_template="llama-3.1", + # ) + + # Load dataset fresh in subprocess + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + alpaca_prompt = """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 = merged_tokenizer.eos_token + + def formatting_prompts_func(examples): + instructions = [] + inputs = [] + outputs = [] + texts = [] + + for conversation in examples["messages"]: + # Extract user message and assistant response + user_message = "" + assistant_message = "" + + for turn in conversation: + if turn["role"] == "user": + user_message = turn["content"] + elif turn["role"] == "assistant": + assistant_message = turn["content"] + + # Store intermediate format + instruction = "Complete the statement" + instructions.append(instruction) + inputs.append(user_message) + outputs.append(assistant_message) + + # Create formatted text + text = alpaca_prompt.format(instruction, user_message, assistant_message) + EOS_TOKEN + texts.append(text) + + return { + "instruction": instructions, + "input": inputs, + "output": outputs, + "text": texts + } + + + + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + # Compute perplexity using the passed dataset + ppl_value = ppl_model(merged_model, merged_tokenizer, dataset_ppl) + + + # IMPORTANT: Convert to Python float if it's a tensor + if torch.is_tensor(ppl_value): + ppl_value = ppl_value.cpu().item() # Move to CPU and convert to Python scalar + elif hasattr(ppl_value, 'item'): + ppl_value = ppl_value.item() # Convert numpy or other array types + else: + ppl_value = float(ppl_value) # Ensure it's a float + + # Return only the perplexity value + result_queue.put(ppl_value) + + # Clean up + del merged_model + del merged_tokenizer + del dataset_ppl + torch.cuda.empty_cache() + gc.collect() + +# Main execution code should be wrapped in this guard +if __name__ == "__main__": + mp.set_start_method('spawn', force=True) + + if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' + else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + + model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/mistral-7b-v0.3", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation + ) + + + EOS_TOKEN = tokenizer.eos_token + + + alpaca_prompt = """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: + {}""" + + + # Define helper functions outside of main + def formatting_prompts_func(examples): + instructions = [] + inputs = [] + outputs = [] + texts = [] + + for conversation in examples["messages"]: + # Extract user message and assistant response + user_message = "" + assistant_message = "" + + for turn in conversation: + if turn["role"] == "user": + user_message = turn["content"] + elif turn["role"] == "assistant": + assistant_message = turn["content"] + + # Store intermediate format + instruction = "Complete the statement" + instructions.append(instruction) + inputs.append(user_message) + outputs.append(assistant_message) + + # Create formatted text + text = alpaca_prompt.format(instruction, user_message, assistant_message) + EOS_TOKEN + texts.append(text) + + + return { + "instruction": instructions, + "input": inputs, + "output": outputs, + "text": texts + } + + + + dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + dataset_train = dataset_train.map(formatting_prompts_func, batched=True) + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + + model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, + ) + + from unsloth import is_bfloat16_supported + + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=200, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), + ) + + # run training + trainer_stats = trainer.train() + + add_to_comparison("Qlora model", ppl_model(model, tokenizer, dataset_ppl)) + + # saving and merging the model to local disk + print("merge and save to local disk") + model.save_pretrained_merged( + save_directory='./unsloth_out/merged_mistral_text_model', + tokenizer=tokenizer + ) + + # print("cleaning") + # del model + # del tokenizer + # torch.cuda.empty_cache() + # gc.collect() + + # load model from local disk and test + print("Loading merged model in 4 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_mistral_text_model", + max_seq_length=2048, + load_in_4bit=True, + load_in_8bit=False, + ) + + add_to_comparison("merged model load 4bit", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + + print("Computing 8-bit model perplexity in subprocess...") + result_queue = mp.Queue() + p = mp.Process(target=load_and_compute_8bit_ppl, args=(result_queue, False, True)) + p.start() + p.join() + + ppl_8bit = result_queue.get() + add_to_comparison("merged model loaded 8bits", ppl_8bit) + + print("Loading merged model in 16 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_mistral_text_model", + max_seq_length=2048, + load_in_4bit=False, + load_in_8bit=False, + ) + + add_to_comparison("merged model loaded 16bits", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + print_model_comparison() + + safe_remove_directory("./outputs") + safe_remove_directory("./unsloth_compiled_cache") + safe_remove_directory("./unsloth_out") diff --git a/tests/saving/language_models/test_merge_model_perplexity_phi_4.py b/tests/saving/language_models/test_merge_model_perplexity_phi_4.py new file mode 100644 index 0000000000..df6f7ac36c --- /dev/null +++ b/tests/saving/language_models/test_merge_model_perplexity_phi_4.py @@ -0,0 +1,222 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc + +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + +# Define helper functions outside of main +def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [ + tokenizer.apply_chat_template( + convo, tokenize = False, add_generation_prompt = False + ) + for convo in convos + ] + return { "text" : texts, } + +def load_and_compute_8bit_ppl(result_queue, load_in_4bit=False, load_in_8bit=False): + """Load model and compute perplexity in subprocess""" + from unsloth import FastLanguageModel + from unsloth.chat_templates import get_chat_template + from tests.utils.perplexity_eval import ppl_model + + # Load model + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_phi4_text_model", + max_seq_length=2048, + load_in_4bit=load_in_4bit, + load_in_8bit=load_in_8bit, + ) + # Set up tokenizer + merged_tokenizer = get_chat_template( + merged_tokenizer, + chat_template="phi-4", + ) + + # Load dataset fresh in subprocess + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + # Format the dataset + def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [merged_tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + # Compute perplexity using the passed dataset + ppl_value = ppl_model(merged_model, merged_tokenizer, dataset_ppl) + + + # IMPORTANT: Convert to Python float if it's a tensor + if torch.is_tensor(ppl_value): + ppl_value = ppl_value.cpu().item() # Move to CPU and convert to Python scalar + elif hasattr(ppl_value, 'item'): + ppl_value = ppl_value.item() # Convert numpy or other array types + else: + ppl_value = float(ppl_value) # Ensure it's a float + + # Return only the perplexity value + result_queue.put(ppl_value) + + # Clean up + del merged_model + del merged_tokenizer + del dataset_ppl + torch.cuda.empty_cache() + gc.collect() + +# Main execution code should be wrapped in this guard +if __name__ == "__main__": + mp.set_start_method('spawn', force=True) + + if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' + else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + + model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/Phi-4", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation + ) + + tokenizer = get_chat_template( + tokenizer, + chat_template="phi-4", + ) + + dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + dataset_train = dataset_train.map(formatting_prompts_func, batched=True) + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + + model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, + ) + + from unsloth import is_bfloat16_supported + + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer), + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=200, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), + ) + + from unsloth.chat_templates import train_on_responses_only + trainer = train_on_responses_only( + trainer, + instruction_part="<|im_start|>user<|im_sep|>\n\n", + response_part="<|im_start|>assistant<|im_sep|>\n\n", + ) + + # run training + trainer_stats = trainer.train() + + add_to_comparison("Qlora model", ppl_model(model, tokenizer, dataset_ppl)) + + # saving and merging the model to local disk + print("merge and save to local disk") + model.save_pretrained_merged( + save_directory='./unsloth_out/merged_phi4_text_model', + tokenizer=tokenizer + ) + + # print("cleaning") + # del model + # del tokenizer + # torch.cuda.empty_cache() + # gc.collect() + + # load model from local disk and test + print("Loading merged model in 4 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_phi4_text_model", + max_seq_length=2048, + load_in_4bit=True, + load_in_8bit=False, + ) + + add_to_comparison("merged model load 4bit", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + + print("Computing 8-bit model perplexity in subprocess...") + result_queue = mp.Queue() + p = mp.Process(target=load_and_compute_8bit_ppl, args=(result_queue, False, True)) + p.start() + p.join() + + ppl_8bit = result_queue.get() + add_to_comparison("merged model loaded 8bits", ppl_8bit) + + print("Loading merged model in 16 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_phi4_text_model", + max_seq_length=2048, + load_in_4bit=False, + load_in_8bit=False, + ) + + add_to_comparison("merged model loaded 16bits", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + print_model_comparison() + + # final cleanup + safe_remove_directory("./outputs") + safe_remove_directory("./unsloth_compiled_cache") + safe_remove_directory("./unsloth_out") diff --git a/tests/saving/language_models/test_merged_model_perplexity_llama-3.1-8b.py b/tests/saving/language_models/test_merged_model_perplexity_llama-3.1-8b.py new file mode 100644 index 0000000000..5338371c04 --- /dev/null +++ b/tests/saving/language_models/test_merged_model_perplexity_llama-3.1-8b.py @@ -0,0 +1,224 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc + +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + +# Define helper functions outside of main +def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + +def load_and_compute_8bit_ppl(result_queue, load_in_4bit=False, load_in_8bit=False): + """Load model and compute perplexity in subprocess""" + from unsloth import FastLanguageModel + from unsloth.chat_templates import get_chat_template + from tests.utils.perplexity_eval import ppl_model + + # Load model + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_llama_text_model", + max_seq_length=2048, + load_in_4bit=load_in_4bit, + load_in_8bit=load_in_8bit, + ) + # Set up tokenizer + merged_tokenizer = get_chat_template( + merged_tokenizer, + chat_template="llama-3.1", + ) + + # Load dataset fresh in subprocess + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + # Format the dataset + def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [merged_tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + # Compute perplexity using the passed dataset + ppl_value = ppl_model(merged_model, merged_tokenizer, dataset_ppl) + + + # IMPORTANT: Convert to Python float if it's a tensor + if torch.is_tensor(ppl_value): + ppl_value = ppl_value.cpu().item() # Move to CPU and convert to Python scalar + elif hasattr(ppl_value, 'item'): + ppl_value = ppl_value.item() # Convert numpy or other array types + else: + ppl_value = float(ppl_value) # Ensure it's a float + + # Return only the perplexity value + result_queue.put(ppl_value) + + # Clean up + del merged_model + del merged_tokenizer + del dataset_ppl + torch.cuda.empty_cache() + gc.collect() + +# Main execution code should be wrapped in this guard +if __name__ == "__main__": + mp.set_start_method('spawn', force=True) + + if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' + else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + + model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/Llama-3.1-8B-Instruct", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation + ) + + tokenizer = get_chat_template( + tokenizer, + chat_template="llama-3.1", + ) + + from unsloth.chat_templates import standardize_sharegpt + dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + dataset_train = dataset_train.map(formatting_prompts_func, batched=True) + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + + print("\n dataset sample [0]") + print(dataset_train[0]) + + add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + + model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, + ) + + from unsloth import is_bfloat16_supported + + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer), + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=200, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), + ) + + from unsloth.chat_templates import train_on_responses_only + trainer = train_on_responses_only( + trainer, + instruction_part="<|start_header_id|>user<|end_header_id|>\n\n", + response_part="<|start_header_id|>assistant<|end_header_id|>\n\n", + ) + + tokenizer.decode(trainer.train_dataset[0]["input_ids"]) + + # run training + trainer_stats = trainer.train() + + add_to_comparison("Qlora model", ppl_model(model, tokenizer, dataset_ppl)) + + # saving and merging the model to local disk + print("merge and save to local disk") + model.save_pretrained_merged( + save_directory='./unsloth_out/merged_llama_text_model', + tokenizer=tokenizer + ) + + # print("cleaning") + # del model + # del tokenizer + # torch.cuda.empty_cache() + # gc.collect() + + # load model from local disk and test + print("Loading merged model in 4 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_llama_text_model", + max_seq_length=2048, + load_in_4bit=True, + load_in_8bit=False, + ) + + add_to_comparison("merged model load 4bit", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + + print("Computing 8-bit model perplexity in subprocess...") + result_queue = mp.Queue() + p = mp.Process(target=load_and_compute_8bit_ppl, args=(result_queue, False, True)) + p.start() + p.join() + + ppl_8bit = result_queue.get() + add_to_comparison("merged model loaded 8bits", ppl_8bit) + + print("Loading merged model in 16 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_llama_text_model", + max_seq_length=2048, + load_in_4bit=False, + load_in_8bit=False, + ) + + add_to_comparison("merged model loaded 16bits", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + print_model_comparison() + + # final cleanup + safe_remove_directory("./outputs") + safe_remove_directory("./unsloth_compiled_cache") + safe_remove_directory("./unsloth_out") diff --git a/tests/saving/language_models/test_merged_model_perplexity_qwen_2.5.py b/tests/saving/language_models/test_merged_model_perplexity_qwen_2.5.py new file mode 100644 index 0000000000..af8e8eaae4 --- /dev/null +++ b/tests/saving/language_models/test_merged_model_perplexity_qwen_2.5.py @@ -0,0 +1,287 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc + +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + + +alpaca_prompt = """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: +{}""" + +# Define helper functions outside of main +def formatting_prompts_func(examples): + instructions = [] + inputs = [] + outputs = [] + texts = [] + + for conversation in examples["messages"]: + # Extract user message and assistant response + user_message = "" + assistant_message = "" + + for turn in conversation: + if turn["role"] == "user": + user_message = turn["content"] + elif turn["role"] == "assistant": + assistant_message = turn["content"] + + # Store intermediate format + instruction = "Complete the statement" + instructions.append(instruction) + inputs.append(user_message) + outputs.append(assistant_message) + + # Create formatted text + text = alpaca_prompt.format(instruction, user_message, assistant_message) + texts.append(text) + + return { + "instruction": instructions, + "input": inputs, + "output": outputs, + "text": texts + } + + +def load_and_compute_8bit_ppl(result_queue, load_in_4bit=False, load_in_8bit=False): + """Load model and compute perplexity in subprocess""" + from unsloth import FastLanguageModel + from tests.utils.perplexity_eval import ppl_model + + # Load model + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_qwen_text_model", + max_seq_length=2048, + load_in_4bit=load_in_4bit, + load_in_8bit=load_in_8bit, + ) + # Set up tokenizer + # merged_tokenizer = get_chat_template( + # merged_tokenizer, + # chat_template="llama-3.1", + # ) + + # Load dataset fresh in subprocess + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + alpaca_prompt = """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: + {}""" + + def formatting_prompts_func(examples): + instructions = [] + inputs = [] + outputs = [] + texts = [] + + for conversation in examples["messages"]: + # Extract user message and assistant response + user_message = "" + assistant_message = "" + + for turn in conversation: + if turn["role"] == "user": + user_message = turn["content"] + elif turn["role"] == "assistant": + assistant_message = turn["content"] + + # Store intermediate format + instruction = "Complete the statement" + instructions.append(instruction) + inputs.append(user_message) + outputs.append(assistant_message) + + # Create formatted text + text = alpaca_prompt.format(instruction, user_message, assistant_message) + texts.append(text) + + return { + "instruction": instructions, + "input": inputs, + "output": outputs, + "text": texts + } + + + + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + # Compute perplexity using the passed dataset + ppl_value = ppl_model(merged_model, merged_tokenizer, dataset_ppl) + + + # IMPORTANT: Convert to Python float if it's a tensor + if torch.is_tensor(ppl_value): + ppl_value = ppl_value.cpu().item() # Move to CPU and convert to Python scalar + elif hasattr(ppl_value, 'item'): + ppl_value = ppl_value.item() # Convert numpy or other array types + else: + ppl_value = float(ppl_value) # Ensure it's a float + + # Return only the perplexity value + result_queue.put(ppl_value) + + # Clean up + # del merged_model + # del merged_tokenizer + # del dataset_ppl + # torch.cuda.empty_cache() + # gc.collect() + +# Main execution code should be wrapped in this guard +if __name__ == "__main__": + mp.set_start_method('spawn', force=True) + + if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' + else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + + model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/Qwen2.5-7B-Instruct", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation + ) + + + dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") + dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + + dataset_train = dataset_train.map(formatting_prompts_func, batched=True) + dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + + add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + + model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, + ) + + from unsloth import is_bfloat16_supported + + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer), + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=200, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), + ) + + # run training + trainer_stats = trainer.train() + + add_to_comparison("Qlora model", ppl_model(model, tokenizer, dataset_ppl)) + + # saving and merging the model to local disk + print("merge and save to local disk") + model.save_pretrained_merged( + save_directory='./unsloth_out/merged_qwen_text_model', + tokenizer=tokenizer + ) + + # print("cleaning") + # del model + # del tokenizer + # torch.cuda.empty_cache() + # gc.collect() + + # load model from local disk and test + print("Loading merged model in 4 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_qwen_text_model", + max_seq_length=2048, + load_in_4bit=True, + load_in_8bit=False, + ) + + add_to_comparison("merged model load 4bit", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + + print("Computing 8-bit model perplexity in subprocess...") + result_queue = mp.Queue() + p = mp.Process(target=load_and_compute_8bit_ppl, args=(result_queue, False, True)) + p.start() + p.join() + + ppl_8bit = result_queue.get() + add_to_comparison("merged model loaded 8bits", ppl_8bit) + + print("Loading merged model in 16 bit for perplexity test") + merged_model, merged_tokenizer = FastLanguageModel.from_pretrained( + model_name="./unsloth_out/merged_qwen_text_model", + max_seq_length=2048, + load_in_4bit=False, + load_in_8bit=False, + ) + + add_to_comparison("merged model loaded 16bits", ppl_model(merged_model, merged_tokenizer, dataset_ppl)) + + print_model_comparison() + + safe_remove_directory("./outputs") + safe_remove_directory("./unsloth_compiled_cache") + safe_remove_directory("./unsloth_out") diff --git a/tests/saving/language_models/test_push_to_hub_merged.py b/tests/saving/language_models/test_push_to_hub_merged.py new file mode 100644 index 0000000000..b77048905b --- /dev/null +++ b/tests/saving/language_models/test_push_to_hub_merged.py @@ -0,0 +1,178 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc +import os +from huggingface_hub import HfFileSystem, hf_hub_download +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + + +# Define helper functions outside of main +def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + + + +if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' +else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + +model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/Llama-3.2-1B-Instruct", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation +) + +tokenizer = get_chat_template( + tokenizer, + chat_template="llama-3.1", +) + +from unsloth.chat_templates import standardize_sharegpt +dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") +dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + +dataset_train = dataset_train.map(formatting_prompts_func, batched=True) +dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + +add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + +model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, +) + +from unsloth import is_bfloat16_supported + +trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer), + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=30, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), +) + +from unsloth.chat_templates import train_on_responses_only +trainer = train_on_responses_only( + trainer, + instruction_part="<|start_header_id|>user<|end_header_id|>\n\n", + response_part="<|start_header_id|>assistant<|end_header_id|>\n\n", +) + +# run training +trainer_stats = trainer.train() + + +# saving and merging the model to local disk +hf_username = os.environ.get("HF_USER", "") +if not hf_username: + hf_username = input("Please enter your Hugging Face username: ").strip() + os.environ["HF_USER"] = hf_username + +hf_token = os.environ.get("HF_TOKEN", "") +if not hf_token: + hf_token = input("Please enter your Hugging Face token: ").strip() + os.environ["HF_TOKEN"] = hf_token + + +repo_name = f"{hf_username}/merged_llama_text_model" +success = { + "upload": False, + "download": False, + } + +# Stage 1: Upload model to Hub +try: + print("\n" + "=" * 80) + print("=== UPLOADING MODEL TO HUB ===".center(80)) + print("=" * 80 + "\n") + model.push_to_hub_merged(repo_name, tokenizer=tokenizer, token=hf_token) + success["upload"] = True + print("✅ Model uploaded successfully!") +except Exception as e: + print(f"❌ Failed to upload model: {e}") + raise Exception("Model upload failed.") + +t +# Stage 2: Test downloading the model (even if cached) +safe_remove_directory(f"./{hf_username}") + +try: + print("\n" + "=" * 80) + print("=== TESTING MODEL DOWNLOAD ===".center(80)) + print("=" * 80 + "\n") + # Force download even if cached + model,tokenizer = FastLanguageModel.from_pretrained(f"{hf_username}/merged_llama_text_model") + success["download"] = True + print("✅ Model downloaded successfully!") +except Exception as e: + print(f"❌ Download failed: {e}") + raise Exception("Model download failed.") + +# Final report +print("\n" + "=" * 80) +print("=== VALIDATION REPORT ===".center(80)) +print("=" * 80 + "\n") +for stage, passed in success.items(): + status = "✓" if passed else "✗" + print(f"{status} {stage.replace('_', ' ').title()}") +print("\n" + "=" * 80) + +if all(success.values()): + print("\n🎉 All stages completed successfully!") +else: + raise Exception("Validation failed for one or more stages.") + +# final cleanup +safe_remove_directory("./outputs") +safe_remove_directory("./unsloth_compiled_cache") diff --git a/tests/saving/language_models/test_push_to_hub_merged_sharded_index_file.py b/tests/saving/language_models/test_push_to_hub_merged_sharded_index_file.py new file mode 100644 index 0000000000..7c31520c45 --- /dev/null +++ b/tests/saving/language_models/test_push_to_hub_merged_sharded_index_file.py @@ -0,0 +1,197 @@ +from unsloth import FastLanguageModel, FastVisionModel, UnslothVisionDataCollator +from unsloth.chat_templates import get_chat_template +from trl import SFTTrainer, SFTConfig +from transformers import DataCollatorForLanguageModeling, DataCollatorForSeq2Seq, TrainingArguments +from datasets import load_dataset, Dataset +import torch +from tqdm import tqdm +import pandas as pd +import multiprocessing as mp +from multiprocessing import Process, Queue +import gc +import os +from huggingface_hub import HfFileSystem, hf_hub_download +# ruff: noqa +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.perplexity_eval import ppl_model, add_to_comparison, print_model_comparison + + +# Define helper functions outside of main +def formatting_prompts_func(examples): + convos = examples["messages"] + texts = [tokenizer.apply_chat_template(convo, tokenize=False, add_generation_prompt=False) for convo in convos] + return {"text": texts} + + + +if torch.cuda.is_bf16_supported(): + compute_dtype = torch.bfloat16 + attn_implementation = 'flash_attention_2' +else: + compute_dtype = torch.float16 + attn_implementation = 'sdpa' + +model, tokenizer = FastLanguageModel.from_pretrained( + model_name="unsloth/Llama-3.1-8B-Instruct", + max_seq_length=2048, + dtype=compute_dtype, + load_in_4bit=True, + load_in_8bit=False, + full_finetuning=False, + attn_implementation=attn_implementation +) + +tokenizer = get_chat_template( + tokenizer, + chat_template="llama-3.1", +) + +from unsloth.chat_templates import standardize_sharegpt +dataset_train = load_dataset("allenai/openassistant-guanaco-reformatted", split="train") +dataset_ppl = load_dataset("allenai/openassistant-guanaco-reformatted", split="eval") + +dataset_train = dataset_train.map(formatting_prompts_func, batched=True) +dataset_ppl = dataset_ppl.map(formatting_prompts_func, batched=True) + +add_to_comparison("Base model 4 bits", ppl_model(model, tokenizer, dataset_ppl)) + +model = FastLanguageModel.get_peft_model( + model, + r=16, + target_modules=['k_proj', 'q_proj', 'v_proj', 'o_proj', "gate_proj", "down_proj", "up_proj"], + lora_alpha=16, + lora_dropout=0, + bias="none", + use_gradient_checkpointing="unsloth", + random_state=3407, + use_rslora=False, + loftq_config=None, +) + +from unsloth import is_bfloat16_supported + +trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + train_dataset=dataset_train, + dataset_text_field="text", + max_seq_length=2048, + data_collator=DataCollatorForSeq2Seq(tokenizer=tokenizer), + dataset_num_proc=2, + packing=False, + args=TrainingArguments( + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + warmup_ratio=0.1, + max_steps=30, + learning_rate=2e-4, + fp16=not is_bfloat16_supported(), + bf16=is_bfloat16_supported(), + logging_steps=50, + optim="adamw_8bit", + lr_scheduler_type="linear", + seed=3407, + output_dir="outputs", + report_to="none", + ), +) + +from unsloth.chat_templates import train_on_responses_only +trainer = train_on_responses_only( + trainer, + instruction_part="<|start_header_id|>user<|end_header_id|>\n\n", + response_part="<|start_header_id|>assistant<|end_header_id|>\n\n", +) + +# run training +trainer_stats = trainer.train() + + +# saving and merging the model to local disk +hf_username = os.environ.get("HF_USER", "") +if not hf_username: + hf_username = input("Please enter your Hugging Face username: ").strip() + os.environ["HF_USER"] = hf_username + +hf_token = os.environ.get("HF_TOKEN", "") +if not hf_token: + hf_token = input("Please enter your Hugging Face token: ").strip() + os.environ["HF_TOKEN"] = hf_token + + +repo_name = f"{hf_username}/merged_llama_text_model" +success = { + "upload": False, + "safetensors_check": False, + "download": False, + } + +# Stage 1: Upload model to Hub +try: + print("\n" + "=" * 80) + print("=== UPLOADING MODEL TO HUB ===".center(80)) + print("=" * 80 + "\n") + model.push_to_hub_merged(repo_name, tokenizer=tokenizer, token=hf_token) + success["upload"] = True + print("✅ Model uploaded successfully!") +except Exception as e: + print(f"❌ Failed to upload model: {e}") + raise Exception("Model upload failed.") + +# Stage 2: Verify safetensors.index.json exists +try: + print("\n" + "=" * 80) + print("=== VERIFYING REPO CONTENTS ===".center(80)) + print("=" * 80 + "\n") + fs = HfFileSystem(token=hf_token) + file_list = fs.ls(repo_name, detail=True) + safetensors_found = any( + file["name"].endswith("model.safetensors.index.json") for file in file_list + ) + if safetensors_found: + success["safetensors_check"] = True + print("✅ model.safetensors.index.json found in repo!") + else: + raise Exception("model.safetensors.index.json not found in repo.") +except Exception as e: + print(f"❌ Verification failed: {e}") + raise Exception("Repo verification failed.") + +# Stage 3: Test downloading the model (even if cached) +safe_remove_directory("./RTannous") + +try: + print("\n" + "=" * 80) + print("=== TESTING MODEL DOWNLOAD ===".center(80)) + print("=" * 80 + "\n") + # Force download even if cached + model,tokenizer = FastLanguageModel.from_pretrained(f"{hf_username}/merged_llama_text_model") + success["download"] = True + print("✅ Model downloaded successfully!") +except Exception as e: + print(f"❌ Download failed: {e}") + raise Exception("Model download failed.") + +# Final report +print("\n" + "=" * 80) +print("=== VALIDATION REPORT ===".center(80)) +print("=" * 80 + "\n") +for stage, passed in success.items(): + status = "✓" if passed else "✗" + print(f"{status} {stage.replace('_', ' ').title()}") +print("\n" + "=" * 80) + +if all(success.values()): + print("\n🎉 All stages completed successfully!") +else: + raise Exception("Validation failed for one or more stages.") + +# final cleanup +safe_remove_directory("./outputs") +safe_remove_directory("./unsloth_compiled_cache") diff --git a/tests/saving/language_models/test_save_merged_grpo_model.py b/tests/saving/language_models/test_save_merged_grpo_model.py new file mode 100644 index 0000000000..35b6f65ccf --- /dev/null +++ b/tests/saving/language_models/test_save_merged_grpo_model.py @@ -0,0 +1,794 @@ +# -*- coding: utf-8 -*- +"""test_Llama3_1_(3B)_GRPO_LoRA (1).ipynb + +### Unsloth + +""" + +from unsloth import FastLanguageModel +import torch +import sys +from pathlib import Path +import multiprocessing as mp +import gc +from multiprocessing import Queue +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.aime_eval import evaluate_model_aime, compare_aime_results + + +max_seq_length = 2048 # Can increase for longer reasoning traces +lora_rank = 64 # Larger rank = smarter, but slower + + +def evaluate_merged_model(result_queue, load_in_4bit=False, load_in_8bit=False): + from unsloth import FastLanguageModel + from tests.utils.aime_eval import evaluate_model_aime + max_seq_length = 2048 # Can increase for longer reasoning traces + lora_rank = 64 # Larger rank = smarter, but slower + + model, tokenizer = FastLanguageModel.from_pretrained( + model_name = "./final_merged_model", + 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.8, # Reduce if out of memory + ) + + print(f"\n{'='*60}") + if load_in_4bit: + print("🔍 EVALUATION Merged model: 4 bits load") + model_type="merged_model_4bits" + elif load_in_8bit: + print("🔍 EVALUATION Merged model: 8 bits load") + model_type="merged_model_8bits" + else: + print("🔍 EVALUATION Merged model: 16 bits load") + model_type="merged_model_16bits" + print(f"{'='*60}") + + evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type=model_type, + temperature=0.3, + n_sampling=8, + max_tokens=32768, + top_p=0.95, + seed=0 + ) + + result_queue.put(results) + + del model + del tokenizer + torch.cuda.empty_cache() + gc.collect() + + + +# Main execution code should be wrapped in this guard +def training_run(result_queue): + model, tokenizer = FastLanguageModel.from_pretrained( + model_name = "meta-llama/Llama-3.2-3B-Instruct", + max_seq_length = max_seq_length, + load_in_4bit = False, # False for LoRA 16bit + fast_inference = True, # Enable vLLM fast inference + max_lora_rank = lora_rank, + gpu_memory_utilization = 0.8, # Reduce if out of memory + ) + + """### Helper Functions + + +#### Helper functions - Data Prep + """ + + import re + import json + + reasoning_start = "" + reasoning_end = "" + solution_start = "" + solution_end = "" + + def extract_hash_answer(text): + """Extract answer from GSM8K format""" + if "####" not in text: + return None + return text.split("####")[1].strip() + + def prepare_gsm8k_dataset(dataset): + """Format GSM8K dataset for training""" + reasoning_start = "" + reasoning_end = "" + solution_start = "" + solution_end = "" + + system_prompt = ( + f"You are given a problem. Think about the problem and reason step by step. " + f"Place your thinking process between {reasoning_start} and {reasoning_end}. " + f"Then, provide your final numerical solution between {solution_start}{solution_end}" + ) + + def format_gsm8k(example): + return { + "prompt": [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": example["question"]}, + ], + "answer": extract_hash_answer(example["answer"]), + } + + return dataset.map(format_gsm8k) + + def prepare_limo_dataset(dataset): + """Format LIMO dataset for SFT training""" + if dataset is None: + return None + + system_prompt = """You are a helpful reasoning assistant. When given a problem, think through it step by step and provide your answer in the following format: + + + [Your detailed step-by-step reasoning and solution process] + + + [Your final numerical answer] + """ + + def format_limo(example): + # Create the assistant response + assistant_response = f"\n{example['solution']}\n\n\n{example['answer']}\n" + + # Return a DICTIONARY with the conversation in a field + return { + "prompt": [ # ← This is the key change - wrap in a dict + {"role": "system", "content": system_prompt}, + {"role": "user", "content": example["question"]}, + {"role": "assistant", "content": assistant_response} + ] + } + + return dataset.map(format_limo) + + print("\n✅ Dataset preparation functions defined!") + + """#### Helper functions - Evaluation""" + + def get_max_prompt_length(dataset, tokenizer): + """Calculate maximum and average prompt length in dataset""" + print("Analyzing prompt lengths...") + + lengths = dataset.map( + lambda x: { + "tokens": tokenizer.apply_chat_template( + x["prompt"], + add_generation_prompt=True, + tokenize=True + ) + }, + batched=True, + ).map(lambda x: {"length": len(x["tokens"])})["length"] + + max_length = max(lengths) + avg_length = sum(lengths) / len(lengths) + min_length = min(lengths) + + print(f"Prompt lengths - Min: {min_length}, Max: {max_length}, Avg: {avg_length:.1f}") + return max_length, avg_length + + def extract_unsloth_answer(text, start_tag="", end_tag=""): + """Extract answer from Unsloth SOLUTION tags""" + pattern = re.escape(start_tag) + r"(.*?)" + re.escape(end_tag) + matches = re.findall(pattern, text, re.DOTALL) + + if matches: + answer = matches[-1] # Get the last match + answer = re.sub(r"[%$,]", "", answer).strip() + return answer + return "" + + def find_number(search_string): + """Find the last number in a string""" + numbers = re.compile( + r"-?[\d,]*\.?\d+", + re.MULTILINE | re.DOTALL | re.IGNORECASE, + ).findall(search_string) + + if numbers: + return numbers[-1].replace(",", "").strip() + return "" + + def remove_symbols(x: str) -> str: + """Remove commas, percent and dollar symbols""" + if not x: + return "" + return x.replace(",", "").replace("%", "").replace("$", "").strip() + + def get_num_tokens(text, tokenizer_instance): + """Count tokens in text""" + if not text: + return 0 + encoding = tokenizer_instance(text, return_tensors="pt") + return len(encoding["input_ids"][0]) + + def check_format_compliance(text, format_type="unsloth"): + """Check if response follows expected format""" + if format_type == "unsloth": + reasoning_start = "" + reasoning_end = "" + solution_start = "" + solution_end = "" + + pattern = ( + rf"^[\s]*{re.escape(reasoning_start)}.+?{re.escape(reasoning_end)}.*?" + rf"{re.escape(solution_start)}.+?{re.escape(solution_end)}[\s]*$" + ) + else: + return False + + return bool(re.match(pattern, text.strip(), re.DOTALL)) + + def normalize_answer(answer): + """Normalize answer for comparison""" + if not answer: + return "" + + normalized = remove_symbols(str(answer)) + + try: + float_val = float(normalized) + if float_val.is_integer(): + return str(int(float_val)) + else: + return str(float_val) + except (ValueError, TypeError): + return normalized + + def evaluate_answer_correctness(extracted_answer, ground_truth): + """Evaluate answer correctness with multiple criteria""" + if not extracted_answer or not ground_truth: + return False, False, 0.0 + + norm_extracted = normalize_answer(extracted_answer) + norm_ground_truth = normalize_answer(ground_truth) + + if norm_extracted == norm_ground_truth: + return True, True, 1.0 + + try: + extracted_num = float(norm_extracted) + ground_truth_num = float(norm_ground_truth) + + if ground_truth_num != 0: + relative_error = abs(extracted_num - ground_truth_num) / abs(ground_truth_num) + + if relative_error < 0.01: + return True, True, 0.9 + elif relative_error < 0.05: + return False, True, 0.7 + elif relative_error < 0.10: + return False, True, 0.5 + else: + if extracted_num == 0: + return True, True, 1.0 + elif abs(extracted_num) < 0.01: + return False, True, 0.7 + + except (ValueError, TypeError): + if norm_extracted.lower() == norm_ground_truth.lower(): + return True, True, 1.0 + + return False, False, 0.0 + + """#### Reward Functions for GRPO""" + + def match_format_exactly(completions, **kwargs): + """Reward function for exact format matching""" + reasoning_start = "" + reasoning_end = "" + solution_start = "" + solution_end = "" + + pattern = ( + rf"^[\s]*{re.escape(reasoning_start)}.+?{re.escape(reasoning_end)}.*?" + rf"{re.escape(solution_start)}.+?{re.escape(solution_end)}[\s]*$" + ) + + responses = [completion[0]["content"] for completion in completions] + rewards = [3.0 if re.match(pattern, response, re.DOTALL) else 0.0 for response in responses] + return rewards + + def match_format_approximately(completions, **kwargs): + """Reward function for approximate format matching""" + reasoning_start = "" + reasoning_end = "" + solution_start = "" + solution_end = "" + + scores = [] + for completion in completions: + score = 0 + response = completion[0]["content"] + score += 0.5 if response.count(reasoning_start) == 1 else -1.0 + score += 0.5 if response.count(reasoning_end) == 1 else -1.0 + score += 0.5 if response.count(solution_start) == 1 else -1.0 + score += 0.5 if response.count(solution_end) == 1 else -1.0 + scores.append(score) + return scores + + def check_answer_correctness(prompts, completions, answer, **kwargs): + """Reward function for answer correctness""" + def extract_solution_answer(text): + pattern = r"(.*?)" + match = re.search(pattern, text, re.DOTALL) + if match: + return re.sub(r"[%$,]", "", match.group(1)).strip() + return "" + + responses = [completion[0]["content"] for completion in completions] + extracted_responses = [extract_solution_answer(r) for r in responses] + + scores = [] + for guess, true_answer in zip(extracted_responses, answer): + score = 0 + if not guess: + scores.append(0) + continue + + if guess == true_answer: + score += 3.0 + elif guess.strip() == true_answer.strip(): + score += 1.5 + else: + try: + ratio = float(guess) / float(true_answer) + if 0.9 <= ratio <= 1.1: + score += 1.0 + elif 0.8 <= ratio <= 1.2: + score += 0.5 + else: + score -= 1.5 + except: + score -= 1.5 + scores.append(score) + return scores + + print("✅ Reward functions defined!") + + """#### Main Evaluation Function""" + + import gc + + + + """#### Comparison and Memory Management""" + + def compare_model_results(all_results): + """Generate comprehensive comparison of multiple model results""" + print(f"\n{'='*80}") + print("COMPREHENSIVE MODEL COMPARISON") + print(f"{'='*80}") + + # Main table + print(f"{'Model':<15} {'Format %':<10} {'Exact %':<10} {'Plausible %':<12} {'Confidence':<12}") + print("-" * 80) + + for result in all_results: + print(f"{result['model_type']:<15} " + f"{result['correct_format_pct']:<10.1f} " + f"{result['exact_match_pct']:<10.1f} " + f"{result['plausible_match_pct']:<12.1f} " + f"{result['avg_confidence']:<12.3f}") + + # Improvement analysis + if len(all_results) > 1: + print(f"\n{'='*50}") + print("IMPROVEMENT ANALYSIS") + print(f"{'='*50}") + + base_result = all_results[0] + for result in all_results[1:]: + print(f"\n{result['model_type']} vs {base_result['model_type']}:") + format_improvement = result['correct_format_pct'] - base_result['correct_format_pct'] + exact_improvement = result['exact_match_pct'] - base_result['exact_match_pct'] + plausible_improvement = result['plausible_match_pct'] - base_result['plausible_match_pct'] + + print(f" Format compliance: {format_improvement:+.1f}%") + print(f" Exact matches: {exact_improvement:+.1f}%") + print(f" Plausible matches: {plausible_improvement:+.1f}%") + + # Save comparison + comparison_data = { + "summary": all_results, + "best_model": max(all_results, key=lambda x: x['exact_match_pct']), + } + + with open("model_comparison_comprehensive.json", "w") as f: + json.dump(comparison_data, f, indent=4) + + print(f"\nBest performing model: {comparison_data['best_model']['model_type']} " + f"({comparison_data['best_model']['exact_match_pct']:.1f}% exact matches)") + + def cleanup_memory(): + """Comprehensive memory cleanup""" + print("🧹 Cleaning up GPU memory...") + for _ in range(10): + torch.cuda.empty_cache() + gc.collect() + + if torch.cuda.is_available(): + allocated = torch.cuda.memory_allocated() / 1024**3 + reserved = torch.cuda.memory_reserved() / 1024**3 + print(f"GPU memory - Allocated: {allocated:.2f} GB, Reserved: {reserved:.2f} GB") + + """#### Data Loading and Preparation""" + + from datasets import load_dataset + + +# Load GSM8K + gsm8k_dataset = load_dataset("openai/gsm8k", "main", split="train") + +# Load LIMO (adjust this based on your access method) + limo_train = load_dataset("GAIR/LIMO", split="train") + +# Prepare datasets + gsm8k_train = prepare_gsm8k_dataset(gsm8k_dataset) + limo_train = prepare_limo_dataset(limo_train) + + + print(f" GSM8K train: {len(gsm8k_train)}") + print(f" LIMO train: {len(limo_train) if limo_train else 0}") + +# Store results + all_results = [] + +# Single temperature evaluation on combined dataset + results = evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type="base", + temperature=0.3, + n_sampling=8, + max_tokens=32768, + top_p=0.95, + seed=0 + ) + + from unsloth.chat_templates import get_chat_template + + tokenizer = get_chat_template( + tokenizer, + chat_template = "llama-3.1", + ) + + def formatting_prompts_func(examples): + convos = examples["prompt"] + texts = [tokenizer.apply_chat_template(convo, tokenize = False, add_generation_prompt = False) for convo in convos] + return { "text" : texts, } + pass + + limo_train = limo_train.map(formatting_prompts_func, batched = True,) + + from trl import SFTTrainer + from transformers import DataCollatorForSeq2Seq, TrainingArguments + from unsloth import is_bfloat16_supported + + + print(f"\n{'*'*60}") + print("🎯 STAGE 1: Qlora Fine-Tuning on LIMO") + print(f"{'*'*60}") + + 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 = 3407, + ) + + + if limo_train is not None: + trainer = SFTTrainer( + model = model, + tokenizer = tokenizer, + train_dataset = limo_train, + dataset_text_field = "text", + max_seq_length = max_seq_length, + data_collator = DataCollatorForSeq2Seq(tokenizer = tokenizer), + dataset_num_proc = 2, + packing = False, # Can make training 5x faster for short sequences. + args = TrainingArguments( + per_device_train_batch_size = 2, + gradient_accumulation_steps = 4, + warmup_steps = 5, + num_train_epochs = 1, # Set this for 1 full training run. + #max_steps = 60, + learning_rate = 2e-4, + fp16 = not is_bfloat16_supported(), + bf16 = is_bfloat16_supported(), + logging_steps = 1, + optim = "adamw_8bit", + weight_decay = 0.01, + lr_scheduler_type = "linear", + seed = 3407, + output_dir = "outputs", + report_to = "none", # Use this for WandB etc + ), + ) + + + from unsloth.chat_templates import train_on_responses_only + trainer = train_on_responses_only( + trainer, + instruction_part = "<|start_header_id|>user<|end_header_id|>\n\n", + response_part = "<|start_header_id|>assistant<|end_header_id|>\n\n", + ) + + # Train + print(f"🚂 Starting SFT training on {len(limo_train)} examples...") + trainer.train() + + # Save checkpoint + model.save_pretrained("qlora_checkpoint") + tokenizer.save_pretrained("qlora_checkpoint") + print("💾 Qlora checkpoint saved!") + + # Cleanup + del trainer + cleanup_memory() + + print("✅ Qlora training completed!") + else: + print("⚠️ Skipping Qlora training - no LIMO dataset available") + +# Cleanup + cleanup_memory() + + global PRINTED_TIMES + PRINTED_TIMES = 0 + global PRINT_EVERY_STEPS + PRINT_EVERY_STEPS = 5 + + match_numbers = re.compile( + solution_start + r".*?([\d\.\,]{1,})", + flags = re.MULTILINE | re.DOTALL + ) + + def check_numbers(prompts, completions, answer, **kwargs): + question = prompts[0][-1]["content"] + responses = [completion[0]["content"] for completion in completions] + + extracted_responses = [ + guess.group(1) + if (guess := match_numbers.search(r)) is not None else None \ + for r in responses + ] + + scores = [] + # Print only every few steps + global PRINTED_TIMES + global PRINT_EVERY_STEPS + if PRINTED_TIMES % PRINT_EVERY_STEPS == 0: + print('*'*20, f"Question:\n{question}", f"\nAnswer:\n{answer[0]}", f"\nResponse:\n{responses[0]}", f"\nExtracted:\n{extracted_responses[0]}") + PRINTED_TIMES += 1 + + for guess, true_answer in zip(extracted_responses, answer): + if guess is None: + scores.append(0) + continue + # Convert to numbers + try: + true_answer = float(true_answer.strip()) + # Remove commas like in 123,456 + guess = float(guess.strip().replace(",", "")) + scores.append(1.5 if guess == true_answer else -0.5) + except: + scores.append(0) + continue + return scores + + print(f"\n{'*'*60}") + print("🎯 STAGE 2: GRPO Fine-Tuning on GSM8K") + print(f"{'*'*60}") + +# Get max prompt length + max_prompt_length, _ = get_max_prompt_length(gsm8k_train, tokenizer) + max_prompt_length = min(max_prompt_length + 10, 512) # Add buffer, cap at 512 + + print(f"Using max_prompt_length: {max_prompt_length}") + + from trl import GRPOConfig, GRPOTrainer + training_args = GRPOConfig( + learning_rate = 5e-6, + weight_decay = 0.1, + warmup_ratio = 0.1, + lr_scheduler_type = "cosine", + optim = "adamw_torch_fused", + logging_steps = 1, + per_device_train_batch_size = 1, + gradient_accumulation_steps = 4, # Increase to 4 for smoother training + num_generations = 8, # Decrease if out of memory + max_prompt_length = max_prompt_length, + max_completion_length = max_seq_length - max_prompt_length, + # num_train_epochs = 1, # Set to 1 for a full training run + #max_steps = 250, + max_steps = 1000, + save_steps = 250, + max_grad_norm = 0.1, + report_to = "none", # Can use Weights & Biases + output_dir = "outputs", + ) + + trainer = GRPOTrainer( + model = model, + processing_class = tokenizer, + reward_funcs = [ + match_format_exactly, + match_format_approximately, + check_answer_correctness, + check_numbers, + ], + args = training_args, + train_dataset = gsm8k_train, + ) + + +# Train + print(f"🚂 Starting GRPO training on {len(gsm8k_train)} examples...") + trainer.train() + +# Save checkpoint + model.save_pretrained("grpo_checkpoint") + tokenizer.save_pretrained("grpo_checkpoint") + print("💾 GRPO checkpoint saved!") + +# Cleanup + del trainer + del training_args + cleanup_memory() + + print("✅ GRPO training completed!") + + print(f"\n{'='*60}") + print("🔍 EVALUATION 3: Final GRPO Model") + print(f"{'='*60}") + + grpo_results = evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type="grpo", + temperature=0.3, + n_sampling=8, + max_tokens=32768, + top_p=0.95, + seed=0 + ) + + all_results.append(grpo_results) + print("✅ Final model evaluation complete!") + + print(f"\n{'='*60}") + print("💾 SAVING FINAL MODEL") + print(f"{'='*60}") + + # Save as merged model + try: + model.save_pretrained_merged("final_merged_model", tokenizer, save_method="merged_16bit") + print("✅ Merged model saved to: final_merged_model/") + except Exception as e: + print(f"⚠️ Could not save merged model: {e}") + print("Final model saved as LoRA adapter only") + + print("💾 Model saving complete!") + + + safe_remove_directory("./unsloth_compiled_cache") + + result_queue.put(results) + + # Clean up + del model + del tokenizer + torch.cuda.empty_cache() + gc.collect() + + + # # Merged model load 16 bits model AIME eval + # result_queue = mp.Queue() + # p = mp.Process(target=evaluate_merged_model, args=(result_queue, False, False)) + # p.start() + # p.join() + # + # merged_16bits = result_queue.get() + # all_results.append(merged_16bits) + # + # # Clean up + # del merged_model + # del merged_tokenizer + # del dataset_ppl + # torch.cuda.empty_cache() + # gc.collect() + # + # safe_remove_directory("./unsloth_compiled_cache") + # + # # Merged model load 8 bits model AIME eval + # + # result_queue = mp.Queue() + # p = mp.Process(target=evaluate_merged_model, args=(result_queue, False, True)) + # p.start() + # p.join() + # + # merged_16bits = result_queue.get() + # all_results.append(merged_16bits) + + + # Merged model load 4 bits AIME eval + # result_queue = mp.Queue() + # p = mp.Process(target=evaluate_merged_model, args=(result_queue, True, False)) + # p.start() + # p.join() + # + # merged_16bits = result_queue.get() + # all_results.append(merged_16bits) + +if __name__ == "__main__": + mp.set_start_method('spawn', force=True) + result_queue = mp.Queue() + all_results = [] + + + # run main finetuning and grpo loop + p = mp.Process(target=training_run, args=(result_queue,)) + p.start() + p.join() + + results = result_queue.get() + all_results = results + + # evaluate merged model loaded 16bits + p = mp.Process(target=evaluate_merged_model, args=(result_queue, False, False)) + p.start() + p.join() + + merged_load_16bits = result_queue.get() + all_results.append(merged_load_16bits) + safe_remove_directory("./unsloth_compiled_cache") + + # Merged model load 8 bits model AIME eval + p = mp.Process(target=evaluate_merged_model, args=(result_queue, False, True)) + p.start() + p.join() + + merged_load_8bits = result_queue.get() + all_results.append(merged_load_8bits) + + safe_remove_directory("./unsloth_compiled_cache") + + # Merged model load 4 bits model AIME eval + p = mp.Process(target=evaluate_merged_model, args=(result_queue, True, False)) + p.start() + p.join() + + merged_load_4bits = result_queue.get() + all_results.append(merged_load_4bits) + + safe_remove_directory("./unsloth_compiled_cache") + + +# AIME-specific comparison function + + print(f"\n{'='*80}") + print("🏆 FINAL TRAINING PIPELINE RESULTS") + print(f"{'='*80}") + +# Use the AIME-specific comparison + compare_aime_results(all_results) diff --git a/tests/saving/vision_models/test_index_file_sharded_model.py b/tests/saving/vision_models/test_index_file_sharded_model.py new file mode 100644 index 0000000000..71dfe24362 --- /dev/null +++ b/tests/saving/vision_models/test_index_file_sharded_model.py @@ -0,0 +1,289 @@ +## Import required libraries + +from unsloth import FastVisionModel, is_bf16_supported +from unsloth.trainer import UnslothVisionDataCollator + +import torch +import os +from datasets import load_dataset +from trl import SFTTrainer, SFTConfig +from huggingface_hub import HfFileSystem +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory + + +## Dataset Preparation""" + +print("\n📊 Loading and preparing dataset...") +dataset = load_dataset("lbourdois/OCR-liboaccn-OPUS-MIT-5M-clean", "en", split="train") +# To select the first 2000 examples +train_dataset = dataset.select(range(2000)) + +# To select the next 200 examples for evaluation +eval_dataset = dataset.select(range(2000, 2200)) + +print(f"✅ Dataset loaded successfully!") +print(f" 📈 Training samples: {len(train_dataset)}") +print(f" 📊 Evaluation samples: {len(eval_dataset)}") +# Convert dataset to OAI messages +def format_data(sample): + return { + "messages": [ + { + "role": "system", + "content": [{"type": "text", "text": system_message}], + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": sample["question"], + }, + { + "type": "image", + "image": sample["image"], + }, + ], + }, + { + "role": "assistant", + "content": [{"type": "text", "text": sample["answer"]}], + }, + ], + } + +print("\n🔄 Formatting dataset for vision training...") +system_message = "You are an expert french ocr system." +# Convert dataset to OAI messages +# need to use list comprehension to keep Pil.Image type, .mape convert image to bytes +train_dataset = [format_data(sample) for sample in train_dataset] +eval_dataset = [format_data(sample) for sample in eval_dataset] +print("✅ Dataset formatting completed!") + +"""## Finetuning Setup and Run""" + + +print("\n" + "=" * 80) +print("=== MODEL LOADING AND SETUP ===".center(80)) +print("=" * 80 + "\n") +# Load Base Model +print("🤖 Loading base vision model...") +try: + model, tokenizer = FastVisionModel.from_pretrained( + # model_name = "unsloth/Qwen2-VL-7B-Instruct", + model_name="unsloth/Qwen2-VL-7B-Instruct", + max_seq_length=2048, # Choose any for long context! + load_in_4bit=True, # 4 bit quantization to reduce memory + load_in_8bit=False, # [NEW!] A bit more accurate, uses 2x memory + full_finetuning=False, # [NEW!] We have full finetuning now! + ) +except Exception as e: + print(f"❌ Failed to load base model: {e}") + raise + +print("\n🔧 Setting up LoRA configuration...") +## Lora Finetuning +try: + model = FastVisionModel.get_peft_model( + model, + finetune_vision_layers=True, # Turn off for just text! + finetune_language_layers=True, # Should leave on! + finetune_attention_modules=True, # Attention good for GRPO + finetune_mlp_modules=True, # SHould leave on always! + r=16, # Choose any number > 0 ! Suggested 8, 16, 32, 64, 128 + lora_alpha=32, + lora_dropout=0, # Supports any, but = 0 is optimized + bias="none", # Supports any, but = "none" is optimized + use_gradient_checkpointing="unsloth", # True or "unsloth" for very long context + random_state=3407, + use_rslora=False, # We support rank stabilized LoRA + loftq_config=None, # And LoftQ + ) + print("✅ LoRA configuration applied successfully!") + print(f" 🎯 LoRA rank (r): 16") + print(f" 📊 LoRA alpha: 32") + print(f" 🔍 Vision layers: Enabled") + print(f" 💬 Language layers: Enabled") +except Exception as e: + print(f"❌ Failed to apply LoRA configuration: {e}") + raise + +print("\n" + "=" * 80) +print("=== TRAINING SETUP ===".center(80)) +print("=" * 80 + "\n") + + +print("🏋️ Preparing trainer...") +FastVisionModel.for_training(model) # Enable for training! + +try: + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + data_collator=UnslothVisionDataCollator(model, tokenizer), + train_dataset=train_dataset, + args=SFTConfig( + # per_device_train_batch_size = 4, + # gradient_accumulation_steps = 8, + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + gradient_checkpointing=True, + gradient_checkpointing_kwargs={ + "use_reentrant": False + }, # use reentrant checkpointing + max_grad_norm=0.3, # max gradient norm based on QLoRA paper + warmup_ratio=0.03, + # num_train_epochs = 2, # Set this instead of max_steps for full training runs + max_steps=10, + learning_rate=2e-4, + fp16=not is_bf16_supported(), + bf16=is_bf16_supported(), + logging_steps=5, + save_strategy="epoch", + optim="adamw_torch_fused", + weight_decay=0.01, + lr_scheduler_type="linear", + seed=3407, + output_dir="checkpoints", + report_to="none", # For Weights and Biases + # You MUST put the below items for vision finetuning: + remove_unused_columns=False, + dataset_text_field="", + dataset_kwargs={"skip_prepare_dataset": True}, + dataset_num_proc=4, + max_seq_length=2048, + ), + ) + print("✅ Trainer setup completed!") + print(f" 📦 Batch size: 2") + print(f" 🔄 Gradient accumulation steps: 4") + print(f" 📈 Max training steps: 10") + print(f" 🎯 Learning rate: 2e-4") + print(f" 💾 Precision: {'BF16' if is_bf16_supported() else 'FP16'}") +except Exception as e: + print(f"❌ Failed to setup trainer: {e}") + raise + +print("\n" + "=" * 80) +print("=== STARTING TRAINING ===".center(80)) +print("=" * 80 + "\n") +# run training +try: + print("🚀 Starting training process...") + trainer_stats = trainer.train() +except Exception as e: + print(f"❌ Training failed: {e}") + raise + +print("\n" + "=" * 80) +print("=== SAVING MODEL ===".center(80)) +print("=" * 80 + "\n") + +print("💾 Saving adapter model and tokenizer locally...") +try: + model.save_pretrained("unsloth-qwen2-7vl-french-ocr-adapter", tokenizer) + tokenizer.save_pretrained("unsloth-qwen2-7vl-french-ocr-adapter") + print("✅ Model saved locally!") +except Exception as e: + print(f"❌ Failed to save model locally: {e}") + raise + + +hf_username = os.environ.get("HF_USER", "") +if not hf_username: + hf_username = input("Please enter your Hugging Face username: ").strip() + os.environ["HF_USER"] = hf_username + +hf_token = os.environ.get("HF_TOKEN", "") +if not hf_token: + hf_token = input("Please enter your Hugging Face token: ").strip() + os.environ["HF_TOKEN"] = hf_token + +repo_name = f"{hf_username}/qwen2-7b-ocr-merged" +success = { + "upload": False, + "safetensors_check": False, + "download": False, +} +# Stage 1: Upload model to Hub +try: + print("\n" + "=" * 80) + print("=== UPLOADING MODEL TO HUB ===".center(80)) + print("=" * 80 + "\n") + print(f"🚀 Uploading to repository: {repo_name}") + model.push_to_hub_merged(repo_name, tokenizer=tokenizer, token=hf_token) + success["upload"] = True + print("✅ Model uploaded successfully!") +except Exception as e: + print(f"❌ Failed to upload model: {e}") + raise Exception("Model upload failed.") + +# Stage 2: Verify safetensors.index.json exists +try: + print("\n" + "=" * 80) + print("=== VERIFYING REPO CONTENTS ===".center(80)) + print("=" * 80 + "\n") + fs = HfFileSystem(token=hf_token) + file_list = fs.ls(repo_name, detail=True) + safetensors_found = any( + file["name"].endswith("model.safetensors.index.json") for file in file_list + ) + if safetensors_found: + success["safetensors_check"] = True + print("✅ model.safetensors.index.json found in repo!") + else: + raise Exception("model.safetensors.index.json not found in repo.") +except Exception as e: + print(f"❌ Verification failed: {e}") + raise Exception("Repo verification failed.") + +# test downloading model even if cached +safe_remove_directory(f"./{hf_username}") + +try: + print("\n" + "=" * 80) + print("=== TESTING MODEL DOWNLOAD ===".center(80)) + print("=" * 80 + "\n") + print("📥 Testing model download...") + # Force download even if cached + test_model, test_tokenizer = FastVisionModel.from_pretrained(repo_name) + success["download"] = True + print("✅ Model downloaded successfully!") + + # Clean up test model + del test_model, test_tokenizer + torch.cuda.empty_cache() +except Exception as e: + print(f"❌ Download failed: {e}") + raise Exception("Model download failed.") + +# Final report +print("\n" + "=" * 80) +print("=== VALIDATION REPORT ===".center(80)) +print("=" * 80 + "\n") +for stage, passed in success.items(): + status = "✅" if passed else "❌" + print(f"{status} {stage.replace('_', ' ').title()}") +print("\n" + "=" * 80) + +if all(success.values()): + print("\n🎉 All stages completed successfully!") + print(f"🌐 Your model is available at: https://huggingface.co/{repo_name}") +else: + raise Exception("Validation failed for one or more stages.") + + +# Final cleanup +print("\n🧹 Cleaning up temporary files...") +safe_remove_directory("./checkpoints") +safe_remove_directory("./unsloth_compiled_cache") +safe_remove_directory("./unsloth-qwen2-7vl-french-ocr-adapter") + +print("\n🎯 Pipeline completed successfully!") +print("=" * 80) diff --git a/tests/saving/vision_models/test_push_to_hub_merged.py b/tests/saving/vision_models/test_push_to_hub_merged.py new file mode 100644 index 0000000000..9845c3e1d0 --- /dev/null +++ b/tests/saving/vision_models/test_push_to_hub_merged.py @@ -0,0 +1,268 @@ +## Import required libraries + +from unsloth import FastVisionModel, is_bf16_supported +from unsloth.trainer import UnslothVisionDataCollator + +import torch +import os +from datasets import load_dataset +from trl import SFTTrainer, SFTConfig + +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory + + +## Dataset Preparation""" + +print("\n📊 Loading and preparing dataset...") +dataset = load_dataset("lbourdois/OCR-liboaccn-OPUS-MIT-5M-clean", "en", split="train") +# To select the first 2000 examples +train_dataset = dataset.select(range(2000)) + +# To select the next 200 examples for evaluation +eval_dataset = dataset.select(range(2000, 2200)) + +print(f"✅ Dataset loaded successfully!") +print(f" 📈 Training samples: {len(train_dataset)}") +print(f" 📊 Evaluation samples: {len(eval_dataset)}") +# Convert dataset to OAI messages +def format_data(sample): + return { + "messages": [ + { + "role": "system", + "content": [{"type": "text", "text": system_message}], + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": sample["question"], + }, + { + "type": "image", + "image": sample["image"], + }, + ], + }, + { + "role": "assistant", + "content": [{"type": "text", "text": sample["answer"]}], + }, + ], + } + +print("\n🔄 Formatting dataset for vision training...") +system_message = "You are an expert french ocr system." +# Convert dataset to OAI messages +# need to use list comprehension to keep Pil.Image type, .mape convert image to bytes +train_dataset = [format_data(sample) for sample in train_dataset] +eval_dataset = [format_data(sample) for sample in eval_dataset] +print("✅ Dataset formatting completed!") + +"""## Finetuning Setup and Run""" + + +print("\n" + "=" * 80) +print("=== MODEL LOADING AND SETUP ===".center(80)) +print("=" * 80 + "\n") +# Load Base Model +print("🤖 Loading base vision model...") +try: + model, tokenizer = FastVisionModel.from_pretrained( + # model_name = "unsloth/Qwen2-VL-7B-Instruct", + model_name="unsloth/Qwen2-VL-2B-Instruct", + max_seq_length=2048, # Choose any for long context! + load_in_4bit=True, # 4 bit quantization to reduce memory + load_in_8bit=False, # [NEW!] A bit more accurate, uses 2x memory + full_finetuning=False, # [NEW!] We have full finetuning now! + ) +except Exception as e: + print(f"❌ Failed to load base model: {e}") + raise + +print("\n🔧 Setting up LoRA configuration...") +## Lora Finetuning +try: + model = FastVisionModel.get_peft_model( + model, + finetune_vision_layers=True, # Turn off for just text! + finetune_language_layers=True, # Should leave on! + finetune_attention_modules=True, # Attention good for GRPO + finetune_mlp_modules=True, # SHould leave on always! + r=16, # Choose any number > 0 ! Suggested 8, 16, 32, 64, 128 + lora_alpha=32, + lora_dropout=0, # Supports any, but = 0 is optimized + bias="none", # Supports any, but = "none" is optimized + use_gradient_checkpointing="unsloth", # True or "unsloth" for very long context + random_state=3407, + use_rslora=False, # We support rank stabilized LoRA + loftq_config=None, # And LoftQ + ) + print("✅ LoRA configuration applied successfully!") + print(f" 🎯 LoRA rank (r): 16") + print(f" 📊 LoRA alpha: 32") + print(f" 🔍 Vision layers: Enabled") + print(f" 💬 Language layers: Enabled") +except Exception as e: + print(f"❌ Failed to apply LoRA configuration: {e}") + raise + +print("\n" + "=" * 80) +print("=== TRAINING SETUP ===".center(80)) +print("=" * 80 + "\n") + + +print("🏋️ Preparing trainer...") +FastVisionModel.for_training(model) # Enable for training! + +try: + trainer = SFTTrainer( + model=model, + tokenizer=tokenizer, + data_collator=UnslothVisionDataCollator(model, tokenizer), + train_dataset=train_dataset, + args=SFTConfig( + # per_device_train_batch_size = 4, + # gradient_accumulation_steps = 8, + per_device_train_batch_size=2, + gradient_accumulation_steps=4, + gradient_checkpointing=True, + gradient_checkpointing_kwargs={ + "use_reentrant": False + }, # use reentrant checkpointing + max_grad_norm=0.3, # max gradient norm based on QLoRA paper + warmup_ratio=0.03, + # num_train_epochs = 2, # Set this instead of max_steps for full training runs + max_steps=10, + learning_rate=2e-4, + fp16=not is_bf16_supported(), + bf16=is_bf16_supported(), + logging_steps=5, + save_strategy="epoch", + optim="adamw_torch_fused", + weight_decay=0.01, + lr_scheduler_type="linear", + seed=3407, + output_dir="checkpoints", + report_to="none", # For Weights and Biases + # You MUST put the below items for vision finetuning: + remove_unused_columns=False, + dataset_text_field="", + dataset_kwargs={"skip_prepare_dataset": True}, + dataset_num_proc=4, + max_seq_length=2048, + ), + ) + print("✅ Trainer setup completed!") + print(f" 📦 Batch size: 2") + print(f" 🔄 Gradient accumulation steps: 4") + print(f" 📈 Max training steps: 10") + print(f" 🎯 Learning rate: 2e-4") + print(f" 💾 Precision: {'BF16' if is_bf16_supported() else 'FP16'}") +except Exception as e: + print(f"❌ Failed to setup trainer: {e}") + raise + +print("\n" + "=" * 80) +print("=== STARTING TRAINING ===".center(80)) +print("=" * 80 + "\n") +# run training +try: + print("🚀 Starting training process...") + trainer_stats = trainer.train() +except Exception as e: + print(f"❌ Training failed: {e}") + raise + +print("\n" + "=" * 80) +print("=== SAVING MODEL ===".center(80)) +print("=" * 80 + "\n") + +print("💾 Saving adapter model and tokenizer locally...") +try: + model.save_pretrained("unsloth-qwen2-7vl-french-ocr-adapter", tokenizer) + tokenizer.save_pretrained("unsloth-qwen2-7vl-french-ocr-adapter") + print("✅ Model saved locally!") +except Exception as e: + print(f"❌ Failed to save model locally: {e}") + raise + + +hf_username = os.environ.get("HF_USER", "") +if not hf_username: + hf_username = input("Please enter your Hugging Face username: ").strip() + os.environ["HF_USER"] = hf_username + +hf_token = os.environ.get("HF_TOKEN", "") +if not hf_token: + hf_token = input("Please enter your Hugging Face token: ").strip() + os.environ["HF_TOKEN"] = hf_token + +repo_name = f"{hf_username}/qwen2-ocr-merged" +success = { + "upload": False, + "download": False, +} +# Stage 1: Upload model to Hub +try: + print("\n" + "=" * 80) + print("=== UPLOADING MODEL TO HUB ===".center(80)) + print("=" * 80 + "\n") + print(f"🚀 Uploading to repository: {repo_name}") + model.push_to_hub_merged(repo_name, tokenizer=tokenizer, token=hf_token) + success["upload"] = True + print("✅ Model uploaded successfully!") +except Exception as e: + print(f"❌ Failed to upload model: {e}") + raise Exception("Model upload failed.") + + +try: + print("\n" + "=" * 80) + print("=== TESTING MODEL DOWNLOAD ===".center(80)) + print("=" * 80 + "\n") + print("📥 Testing model download...") + # Force download even if cached + test_model, test_tokenizer = FastVisionModel.from_pretrained(repo_name) + success["download"] = True + print("✅ Model downloaded successfully!") + + # Clean up test model + del test_model, test_tokenizer + torch.cuda.empty_cache() +except Exception as e: + print(f"❌ Download failed: {e}") + raise Exception("Model download failed.") + +# Final report +print("\n" + "=" * 80) +print("=== VALIDATION REPORT ===".center(80)) +print("=" * 80 + "\n") +for stage, passed in success.items(): + status = "✅" if passed else "❌" + print(f"{status} {stage.replace('_', ' ').title()}") +print("\n" + "=" * 80) + +if all(success.values()): + print("\n🎉 All stages completed successfully!") + print(f"🌐 Your model is available at: https://huggingface.co/{repo_name}") +else: + raise Exception("Validation failed for one or more stages.") + + +# Final cleanup +print("\n🧹 Cleaning up temporary files...") +safe_remove_directory("./checkpoints") +safe_remove_directory("./unsloth_compiled_cache") +safe_remove_directory("./unsloth-qwen2-7vl-french-ocr-adapter") +safe_remove_directory(f"./{hf_username}") + +print("\n🎯 Pipeline completed successfully!") +print("=" * 80) diff --git a/tests/saving/vision_models/test_save_merge_vision_model_ocr_benchmark.py b/tests/saving/vision_models/test_save_merge_vision_model_ocr_benchmark.py new file mode 100644 index 0000000000..8f2617e556 --- /dev/null +++ b/tests/saving/vision_models/test_save_merge_vision_model_ocr_benchmark.py @@ -0,0 +1,254 @@ +# -*- coding: utf-8 -*- + +from unsloth import FastVisionModel + +import torch +from qwen_vl_utils import process_vision_info +import os +from datasets import load_dataset +from trl import SFTTrainer, SFTConfig + +import sys +from pathlib import Path + +REPO_ROOT = Path(__file__).parents[3] +sys.path.append(str(REPO_ROOT)) + +from tests.utils.cleanup_utils import safe_remove_directory +from tests.utils.ocr_eval import OCRModelEvaluator + + +## Dataset Preparation +from datasets import load_dataset + +dataset = load_dataset("lbourdois/OCR-liboaccn-OPUS-MIT-5M-clean", 'en', split="train") +# To select the first 2000 examples +train_dataset = dataset.select(range(2000)) + +# To select the next 200 examples for evaluation +eval_dataset = dataset.select(range(2000, 2200)) + +# Convert dataset to OAI messages +def format_data(sample): + return {"messages": [ + { + "role": "system", + "content": [{"type": "text", "text": system_message}], + }, + { + "role": "user", + "content": [ + { + "type": "text", + "text": sample["question"], + },{ + "type": "image", + "image": sample["image"], + } + ], + }, + { + "role": "assistant", + "content": [{"type": "text", "text": sample["answer"]}], + }, + ], + } + +system_message = "You are an expert french ocr system." +# Convert dataset to OAI messages +# need to use list comprehension to keep Pil.Image type, .mape convert image to bytes +train_dataset = [format_data(sample) for sample in train_dataset] +eval_dataset = [format_data(sample) for sample in eval_dataset] + +## Setup OCR main evaluation function and helpers +import os +import torch +from tqdm import tqdm +import pandas as pd +from jiwer import wer, cer +from qwen_vl_utils import process_vision_info + +# +ocr_evaluator = OCRModelEvaluator() +model_comparison_results = {} + +## Finetuning Setup and Run +# Load Base Model + +model, tokenizer = FastVisionModel.from_pretrained( + model_name = "unsloth/Qwen2-VL-7B-Instruct", + max_seq_length = 2048, # Choose any for long context! + load_in_4bit = True, # 4 bit quantization to reduce memory + load_in_8bit = False, # [NEW!] A bit more accurate, uses 2x memory + full_finetuning = False, # [NEW!] We have full finetuning now! +) + +# benchmark base model performance +model_name = "Unsloth Base model" +FastVisionModel.for_inference(model) +avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_base_model_results") +ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) + +## Lora Finetuning +model = FastVisionModel.get_peft_model( + model, + finetune_vision_layers = True, # Turn off for just text! + finetune_language_layers = True, # Should leave on! + finetune_attention_modules = True, # Attention good for GRPO + finetune_mlp_modules = True, # SHould leave on always! + + r = 16, # 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",], + lora_alpha = 32, + lora_dropout = 0, # Supports any, but = 0 is optimized + bias = "none", # Supports any, but = "none" is optimized + # [NEW] "unsloth" uses 30% less VRAM, fits 2x larger batch sizes! + use_gradient_checkpointing = "unsloth", # True or "unsloth" for very long context + random_state = 3407, + use_rslora = False, # We support rank stabilized LoRA + loftq_config = None, # And LoftQ +) + +from unsloth import is_bf16_supported +from unsloth.trainer import UnslothVisionDataCollator +FastVisionModel.for_training(model) # Enable for training! +model.config.use_cache = False + + +trainer = SFTTrainer( + model = model, + tokenizer = tokenizer, + data_collator = UnslothVisionDataCollator(model, tokenizer), + train_dataset = train_dataset, + args = SFTConfig( + #per_device_train_batch_size = 4, + #gradient_accumulation_steps = 8, + per_device_train_batch_size = 2, + gradient_accumulation_steps = 4, + gradient_checkpointing=True, + gradient_checkpointing_kwargs = {"use_reentrant": False}, # use reentrant checkpointing + max_grad_norm=0.3, # max gradient norm based on QLoRA paper + warmup_ratio=0.03, + #num_train_epochs = 2, # Set this instead of max_steps for full training runs + max_steps=60, + learning_rate = 2e-4, + fp16 = not is_bf16_supported(), + bf16 = is_bf16_supported(), + logging_steps = 5, + save_strategy="epoch", + optim = "adamw_torch_fused", + weight_decay = 0.01, + lr_scheduler_type = "linear", + seed = 3407, + output_dir = "unsloth-qwen2-7vl-french-ocr-checkpoints", + report_to = "none", # For Weights and Biases + + # You MUST put the below items for vision finetuning: + remove_unused_columns = False, + dataset_text_field = "", + dataset_kwargs = {"skip_prepare_dataset": True}, + dataset_num_proc = 4, + max_seq_length = 2048, + ), +) + +# run training +trainer_stats = trainer.train() + +model.save_pretrained("unsloth-qwen2-7vl-french-ocr-adapter", tokenizer) +tokenizer.save_pretrained("unsloth-qwen2-7vl-french-ocr-adapter") + +## Measure Adapter Performance + +# benchmark lora model performance +model_name = "Unsloth lora adapter model" +FastVisionModel.for_inference(model) +avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_lora_model_results") +ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) + +## Merge Model + +def find_lora_base_model(model_to_inspect): + current = model_to_inspect + if hasattr(current, "base_model"): + current = current.base_model + if hasattr(current, "model"): + current = current.model + return current +pass + +base = find_lora_base_model(model) + +print((base.__class__.__name__)) + +# merge default 16 bits +model.save_pretrained_merged(save_directory="qwen2-ocr-merged-finetune-merge-16bit", tokenizer=tokenizer) + + +## Benchmark merged model performance + +### 16 bits merged model + +model, tokenizer = FastVisionModel.from_pretrained("./qwen2-ocr-merged-finetune-merge-16bit",load_in_4bit=False, load_in_8bit=False) + +# benchmark 4bit loaded, 16bits merged model performance +model_name = "Unsloth 16bits-merged model load-16bits" +model.config.use_cache = True + +avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_16bits_merged_model_load_16bits_results") +ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) + +# load 16bits-merged model in 4 bits +model, tokenizer = FastVisionModel.from_pretrained("./qwen2-ocr-merged-finetune-merge-16bit",load_in_4bit=True, load_in_8bit=False) + +# benchmark 4bit loaded, 16bits merged model performance +model_name = "Unsloth 16bits-merged model load-4bits" +model.config.use_cache = True + +avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_16bits_merged_model_load_4bits_results") +ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) + +# load model in 8 bits +model, tokenizer = FastVisionModel.from_pretrained("./qwen2-ocr-merged-finetune-merge-16bit",load_in_4bit=False, load_in_8bit=True) + +# benchmark 4bit loaded, 16bits merged model performance +model_name = "Unsloth 16bits-merged model load-8bits" +avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_16bits_merged_model_load_8bits_results") +ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) + +# """### 4 bits merged model""" +# +# # load 4bits-merged model in 4 bits +# model, tokenizer = FastVisionModel.from_pretrained("./qwen2-ocr-merged-finetune-merge-4bit",load_in_4bit=True, load_in_8bit=False) +# +# # benchmark 4bit loaded, 4bits merged model performance +# model_name = "Unsloth 4bits-merged model load-4bits" +# +# avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_4bits_merged_model_load_4bits_results") +# ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) +# +# # load model in 8 bits +# model, tokenizer = FastVisionModel.from_pretrained("./qwen2-ocr-merged-finetune-merge-4bit",load_in_4bit=False, load_in_8bit=True) +# +# # benchmark 8bit loaded, 4bits merged model performance +# model_name = "Unsloth 4bits-merged model load-8bits" +# +# avg_wer, avg_cer = ocr_evaluator.evaluate_model(model, tokenizer, eval_dataset, output_dir="unsloth_4bits_merged_model_load_8bits_results") +# ocr_evaluator.add_to_comparison(model_name, avg_wer, avg_cer) + +# Model comparison report +#print model comparison +ocr_evaluator.print_model_comparison() + + + +# Final cleanup +print("\n🧹 Cleaning up temporary files...") +safe_remove_directory("./unsloth-qwen2-7vl-french-ocr-adapter") +safe_remove_directory("./unsloth-qwen2-7vl-french-ocr-checkpoints") +safe_remove_directory("./unsloth_compiled_cache") +safe_remove_directory("./qwen2-ocr-merged-finetune-merge-16bit") + +print("\n🎯 Pipeline completed successfully!") +print("=" * 80) diff --git a/tests/utils/aime_eval.md b/tests/utils/aime_eval.md new file mode 100644 index 0000000000..217e939094 --- /dev/null +++ b/tests/utils/aime_eval.md @@ -0,0 +1,264 @@ +# AIME Dataset Evaluator + +A Python module for evaluating language models on the AIME (American Invitational Mathematics Examination) dataset. This evaluator automatically downloads and combines multiple AIME test datasets and provides comprehensive mathematical reasoning assessment. + + +## Basic Usage + +```python +from aime_utils import evaluate_model_aime + +# Simple AIME evaluation +results = evaluate_model_aime( + model=your_model, + tokenizer=your_tokenizer, + model_type="base_model", + temperature=0.3, + n_sampling=8, + max_tokens=32768 +) + +print(f"AIME Accuracy: {results['accuracy']:.1f}%") +print(f"Pass@8: {results['pass_at_k']:.1f}%") +``` + +## Advanced Usage + +```python +from aime_utils import evaluate_model_aime, compare_aime_results + +# Evaluate multiple model configurations +all_results = [] + +# Base model +base_results = evaluate_model_aime( + model=base_model, + tokenizer=tokenizer, + model_type="base", + temperature=0.3, + n_sampling=8 +) +all_results.append(base_results) + +# Fine-tuned model +ft_results = evaluate_model_aime( + model=finetuned_model, + tokenizer=tokenizer, + model_type="finetuned", + temperature=0.3, + n_sampling=8 +) +all_results.append(ft_results) + +# Generate comprehensive comparison +compare_aime_results(all_results) +``` + +## Dataset Format + +The evaluator automatically handles AIME dataset format with problems containing: + +- **Problem**: Mathematical question text +- **Answer**: Numerical answer (0-999 range for AIME) +- **Solution**: Step-by-step solution (when available) +- **Source**: Original dataset identifier (test2024, test2025-I, test2025-II) + +```python +# Automatic dataset download and formatting +{ + "global_id": 0, + "original_id": "problem_1", + "source_dataset": "test2024", + "problem": "Find the number of...", + "answer": "123", + "solution": "Step-by-step solution...", + "prompt": [ + {"role": "system", "content": "You are a mathematical problem solver..."}, + {"role": "user", "content": "Problem: Find the number of..."} + ] +} +``` + + +## Configuration Examples + +### Conservative Evaluation +```python +# Lower temperature for more consistent answers +results = evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type="conservative", + temperature=0.1, + n_sampling=4, + top_p=0.9 +) +``` + +### High-Sample Evaluation +```python +# More samples for better Pass@K estimation +results = evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type="high_sample", + temperature=0.5, + n_sampling=16, + max_tokens=16384 +) +``` + +### Memory-Optimized +```python +# Reduced parameters for limited resources +results = evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type="lite", + temperature=0.3, + n_sampling=4, + max_tokens=8192 +) +``` + +## Examples + +### Complete Model Pipeline Evaluation +```python +from aime_utils import evaluate_model_aime, compare_aime_results + +def evaluate_training_pipeline(base_model, finetuned_model, merged_model, tokenizer): + """Evaluate complete training pipeline on AIME""" + + all_results = [] + + # Standard evaluation configuration + eval_config = { + "temperature": 0.3, + "n_sampling": 8, + "max_tokens": 32768, + "top_p": 0.95, + "seed": 0 + } + + # Evaluate base model + print("Evaluating base model...") + base_results = evaluate_model_aime( + model=base_model, + tokenizer=tokenizer, + model_type="base", + **eval_config + ) + all_results.append(base_results) + + # Evaluate fine-tuned model + print("Evaluating fine-tuned model...") + ft_results = evaluate_model_aime( + model=finetuned_model, + tokenizer=tokenizer, + model_type="finetuned", + **eval_config + ) + all_results.append(ft_results) + + # Evaluate merged model + print("Evaluating merged model...") + merged_results = evaluate_model_aime( + model=merged_model, + tokenizer=tokenizer, + model_type="merged", + **eval_config + ) + all_results.append(merged_results) + + # Generate comparison report + compare_aime_results(all_results) + + return all_results +``` + +### Quantization Impact Analysis +```python +def analyze_quantization_impact(model_paths, tokenizer): + """Analyze impact of different quantization levels""" + + quantization_configs = { + "fp16": {"load_in_4bit": False, "load_in_8bit": False}, + "8bit": {"load_in_4bit": False, "load_in_8bit": True}, + "4bit": {"load_in_4bit": True, "load_in_8bit": False} + } + + all_results = [] + + for quant_name, load_config in quantization_configs.items(): + print(f"Evaluating {quant_name} quantization...") + + # Load model with specific quantization + model = load_model_with_config(model_paths["merged"], **load_config) + + results = evaluate_model_aime( + model=model, + tokenizer=tokenizer, + model_type=f"merged_{quant_name}", + temperature=0.3, + n_sampling=8, + max_tokens=32768 + ) + all_results.append(results) + + # Cleanup + del model + torch.cuda.empty_cache() + + compare_aime_results(all_results) + return all_results +``` + +## Output Format + +### Individual Evaluation Results +``` +🧮 AIME EVALUATION - BASE MODEL +Combined Dataset: test2024 + test2025-I + test2025-II +==================================================================== + +🎯 Overall Performance: + Total problems: 45 + Correct answers: 12/45 (26.7%) + Pass@8: 31.1% + +📈 Performance by Dataset: + test2024: 4/15 (26.7%) + test2025-I: 5/15 (33.3%) + test2025-II: 3/15 (20.0%) + +🎖️ AIME Performance: ✅ EXCELLENT (26.7%) +``` + +### Comparison Report +``` +COMPREHENSIVE AIME MODEL COMPARISON +================================================================================ +Model Accuracy % Pass@K % Correct Total +-------------------------------------------------------------------------------- +finetuned 31.1 35.6 14 45 +base 26.7 31.1 12 45 +merged_4bit 24.4 28.9 11 45 + +IMPROVEMENT ANALYSIS +================================================== +finetuned vs base: + Accuracy improvement: +4.4% + Pass@K improvement: +4.5% +``` + +## Performance Tiers + +The evaluator provides performance assessment based on AIME difficulty: + +- **🏆 EXCEPTIONAL**: ≥50% accuracy +- **✅ EXCELLENT**: ≥30% accuracy +- **🎯 VERY GOOD**: ≥20% accuracy +- **⚠️ GOOD**: ≥10% accuracy +- **📈 FAIR**: ≥5% accuracy +- **❌ NEEDS IMPROVEMENT**: <5% accuracy diff --git a/tests/utils/aime_eval.py b/tests/utils/aime_eval.py new file mode 100644 index 0000000000..54b0d8e511 --- /dev/null +++ b/tests/utils/aime_eval.py @@ -0,0 +1,495 @@ +""" +AIME Dataset Evaluation Module + +This module provides functions to evaluate language models on the combined AIME dataset +(test2024 + test2025-I + test2025-II). +""" + +import json +import requests +import os +import re +import logging +from typing import List, Dict, Any +from tqdm import tqdm +from vllm import SamplingParams + + +def download_and_combine_aime_datasets(data_dir: str = "./data/aime") -> str: + """Download all AIME datasets and combine them into a single file""" + + datasets = { + "test2024": "https://raw.githubusercontent.com/GAIR-NLP/AIME-Preview/main/eval/data/aime/test2024.jsonl", + "test2025-I": "https://raw.githubusercontent.com/GAIR-NLP/AIME-Preview/main/eval/data/aime/test2025-I.jsonl", + "test2025-II": "https://raw.githubusercontent.com/GAIR-NLP/AIME-Preview/main/eval/data/aime/test2025-II.jsonl" + } + + os.makedirs(data_dir, exist_ok=True) + combined_filepath = os.path.join(data_dir, "aime.jsonl") + + # Check if combined file already exists + if os.path.exists(combined_filepath): + print(f"Combined AIME dataset already exists at {combined_filepath}") + return combined_filepath + + print("Downloading and combining AIME datasets...") + + all_problems = [] + global_id = 0 + + for dataset_name, url in datasets.items(): + print(f" Downloading {dataset_name}...") + + try: + response = requests.get(url) + response.raise_for_status() + + # Parse each line and add source information + for line_num, line in enumerate(response.text.strip().split('\n')): + if line.strip(): + try: + data = json.loads(line) + # Add source dataset information and global ID + data['source_dataset'] = dataset_name + data['original_id'] = data.get('id', line_num) + data['global_id'] = global_id + global_id += 1 + all_problems.append(data) + except json.JSONDecodeError as e: + print(f" Warning: Error parsing line {line_num + 1} in {dataset_name}: {e}") + continue + + except requests.RequestException as e: + print(f" Error downloading {dataset_name}: {e}") + continue + + # Write combined dataset + if all_problems: + with open(combined_filepath, 'w', encoding='utf-8') as f: + for problem in all_problems: + f.write(json.dumps(problem, ensure_ascii=False) + '\n') + + print(f"✅ Combined {len(all_problems)} problems from {len(datasets)} datasets") + print(f" Saved to: {combined_filepath}") + + # Print summary by dataset + for dataset_name in datasets.keys(): + count = sum(1 for p in all_problems if p['source_dataset'] == dataset_name) + print(f" {dataset_name}: {count} problems") + + else: + raise RuntimeError("No problems were successfully downloaded") + + return combined_filepath + + +def load_aime_dataset(data_dir: str = "./data/aime") -> List[Dict[str, Any]]: + """Load combined AIME dataset and format for evaluation""" + + # Download and combine if needed + filepath = download_and_combine_aime_datasets(data_dir) + + examples = [] + with open(filepath, 'r', encoding='utf-8') as f: + for line_num, line in enumerate(f): + line = line.strip() + if line: + try: + data = json.loads(line) + + # Format as expected by our evaluation + formatted_example = { + "global_id": data.get("global_id", line_num), + "original_id": data.get("original_id", data.get("id", line_num)), + "source_dataset": data.get("source_dataset", "unknown"), + "problem": data["problem"], + "answer": str(data["answer"]), # Ensure answer is string + "solution": data.get("solution", ""), + "url": data.get("url", ""), + # Format as chat messages for the model + "prompt": [ + {"role": "system", "content": "You are a mathematical problem solver. Solve the given problem step by step and provide your final answer clearly."}, + {"role": "user", "content": f"Problem: {data['problem']}\n\nSolve this step by step and provide your final numerical answer."} + ] + } + examples.append(formatted_example) + + except json.JSONDecodeError as e: + print(f"Error parsing line {line_num + 1}: {e}") + continue + + print(f"Loaded {len(examples)} problems from combined AIME dataset") + + # Print breakdown by source + source_counts = {} + for example in examples: + source = example['source_dataset'] + source_counts[source] = source_counts.get(source, 0) + 1 + + for source, count in source_counts.items(): + print(f" {source}: {count} problems") + + return examples + + +def extract_aime_answer(response: str) -> str: + """Extract numerical answer from AIME response""" + + # AIME answers are integers from 0-999 + # Look for patterns like "The answer is 123" or just standalone numbers + patterns = [ + r"(?:the )?(?:final )?answer is (\d{1,3})", + r"(?:therefore|thus|so),?\s*(?:the )?(?:final )?answer is (\d{1,3})", + r"\\boxed\{(\d{1,3})\}", + r"\$\\boxed\{(\d{1,3})\}\$", + r"(?:answer|result):\s*(\d{1,3})", + r"(?:^|\n)\s*(\d{1,3})\s*(?:\n|$)", # Standalone number + ] + + response_lower = response.lower().strip() + + for pattern in patterns: + matches = re.findall(pattern, response_lower, re.MULTILINE | re.IGNORECASE) + if matches: + # Get the last match (most likely to be final answer) + answer = matches[-1] + try: + num = int(answer) + if 0 <= num <= 999: # AIME answers are in range 0-999 + return str(num) + except ValueError: + continue + + # If no clear pattern found, try to extract any 1-3 digit number + numbers = re.findall(r'\b(\d{1,3})\b', response) + if numbers: + for num_str in reversed(numbers): # Check from end + try: + num = int(num_str) + if 0 <= num <= 999: + return str(num) + except ValueError: + continue + + return "" + + +def get_num_tokens(text, tokenizer_instance): + """Count tokens in text""" + if not text: + return 0 + encoding = tokenizer_instance(text, return_tensors="pt") + return len(encoding["input_ids"][0]) + + +def evaluate_model_aime(model, tokenizer, model_type="base", lora_request=None, + temperature=0.3, n_sampling=8, max_tokens=32768, top_p=0.95, seed=0): + """Evaluate model on combined AIME dataset with official configuration""" + + print(f"\n{'='*70}") + print(f"🧮 AIME EVALUATION - {model_type.upper()} MODEL") + print(f"Combined Dataset: test2024 + test2025-I + test2025-II") + print(f"{'='*70}") + + # Load combined AIME dataset + try: + eval_dataset = load_aime_dataset() + except Exception as e: + print(f"Error loading dataset: {e}") + return None + + if not eval_dataset: + print("No examples found in dataset") + return None + + # Initialize tracking variables + records = {} + input_tokens = [] + output_tokens = [] + correct_answers = 0 + + # Track performance by source dataset + source_stats = {} + for example in eval_dataset: + source = example['source_dataset'] + if source not in source_stats: + source_stats[source] = {'total': 0, 'correct': 0} + source_stats[source]['total'] += 1 + + # Setup sampling parameters (AIME configuration) + sampling_params = SamplingParams( + temperature=temperature, + top_p=top_p, + max_tokens=max_tokens, + n=n_sampling, # Multiple samples per question + seed=seed, + ) + + print(f"\n🔧 Configuration:") + print(f" Temperature: {temperature}") + print(f" Samples per question: {n_sampling}") + print(f" Max tokens: {max_tokens}") + print(f" Top-p: {top_p}") + print(f" Seed: {seed}") + + # Temporarily suppress verbose logging + original_levels = {} + loggers_to_suppress = ['vllm', 'vllm.engine', 'vllm.worker', 'vllm.model_executor', 'vllm.executor', 'ray'] + + for logger_name in loggers_to_suppress: + logger = logging.getLogger(logger_name) + original_levels[logger_name] = logger.level + logger.setLevel(logging.WARNING) + + try: + print(f"\n🚀 Evaluating {len(eval_dataset)} problems...") + + # Main evaluation loop + with tqdm(total=len(eval_dataset), desc="Processing AIME problems", unit="problem") as pbar: + for task_id, item in enumerate(eval_dataset): + try: + # Prepare prompt + prompt_text = tokenizer.apply_chat_template( + item["prompt"], + add_generation_prompt=True, + tokenize=False + ) + + input_tokens.append(get_num_tokens(prompt_text, tokenizer)) + + # Generate multiple responses + outputs = model.fast_generate( + [prompt_text], + sampling_params=sampling_params, + lora_request=lora_request, + use_tqdm=False, + )[0].outputs + + # Process all generated responses + responses = [output.text for output in outputs] + extracted_answers = [extract_aime_answer(response) for response in responses] + + # Calculate total output tokens + total_output_tokens = sum(get_num_tokens(response, tokenizer) for response in responses) + output_tokens.append(total_output_tokens) + + # Check if any answer is correct + ground_truth = item["answer"] + correct_responses = [ans == ground_truth for ans in extracted_answers] + is_correct = any(correct_responses) + + if is_correct: + correct_answers += 1 + source_stats[item['source_dataset']]['correct'] += 1 + + # Store detailed record + records[task_id] = { + "global_id": item["global_id"], + "original_id": item["original_id"], + "source_dataset": item["source_dataset"], + "problem": item["problem"], + "ground_truth": ground_truth, + "responses": responses, + "extracted_answers": extracted_answers, + "correct_responses": correct_responses, + "is_correct": is_correct, + "input_tokens": input_tokens[-1], + "output_tokens": total_output_tokens, + "n_correct": sum(correct_responses), + "n_total": len(responses), + "solution": item.get("solution", ""), + "url": item.get("url", "") + } + + # Update progress + current_accuracy = correct_answers / (task_id + 1) * 100 + pbar.set_postfix({ + 'accuracy': f'{current_accuracy:.1f}%', + 'correct': correct_answers, + 'total': task_id + 1 + }) + pbar.update(1) + + except Exception as e: + print(f"\nError processing problem {task_id}: {str(e)}") + records[task_id] = { + "global_id": item.get("global_id", task_id), + "original_id": item.get("original_id", task_id), + "source_dataset": item.get("source_dataset", "unknown"), + "problem": item["problem"], + "ground_truth": item["answer"], + "error": str(e), + "is_correct": False + } + pbar.update(1) + continue + + finally: + # Restore logging levels + for logger_name, level in original_levels.items(): + logging.getLogger(logger_name).setLevel(level) + + # Calculate metrics + total_problems = len(eval_dataset) + accuracy = correct_answers / total_problems * 100 + + # Calculate Pass@k (probability that at least one of k samples is correct) + pass_at_k_scores = [] + for record in records.values(): + if "n_correct" in record and "n_total" in record: + n_correct = record["n_correct"] + n_total = record["n_total"] + if n_correct > 0: + pass_at_k_scores.append(1.0) + else: + pass_at_k_scores.append(0.0) + + pass_at_k = sum(pass_at_k_scores) / len(pass_at_k_scores) if pass_at_k_scores else 0 + + # Calculate per-source accuracies + source_accuracies = {} + for source, stats in source_stats.items(): + source_accuracies[source] = (stats['correct'] / stats['total'] * 100) if stats['total'] > 0 else 0 + + results = { + "model_type": model_type, + "dataset": "aime_combined", + "total_problems": total_problems, + "correct_answers": correct_answers, + "accuracy": accuracy, + "pass_at_k": pass_at_k * 100, + "source_stats": source_stats, + "source_accuracies": source_accuracies, + "temperature": temperature, + "n_sampling": n_sampling, + "max_tokens": max_tokens, + "top_p": top_p, + "seed": seed, + "avg_input_tokens": sum(input_tokens) / len(input_tokens) if input_tokens else 0, + "avg_output_tokens": sum(output_tokens) / len(output_tokens) if output_tokens else 0, + "max_input_tokens": max(input_tokens) if input_tokens else 0, + "max_output_tokens": max(output_tokens) if output_tokens else 0, + } + + # Save results + filename = f"aime_eval_combined_{model_type}_t{temperature}_n{n_sampling}.json" + with open(filename, "w", encoding="utf-8") as f: + json.dump({"results": results, "records": records}, f, indent=4) + + # Print comprehensive summary + print(f"\n{'='*70}") + print(f"📊 AIME EVALUATION RESULTS - {model_type.upper()}") + print(f"{'='*70}") + + print(f"\n🎯 Overall Performance:") + print(f" Total problems: {total_problems:>6}") + print(f" Correct answers: {correct_answers:>6}/{total_problems} ({accuracy:>5.1f}%)") + print(f" Pass@{n_sampling}: {pass_at_k:>10.1f}%") + + print(f"\n📈 Performance by Dataset:") + for source, stats in source_stats.items(): + source_acc = source_accuracies[source] + print(f" {source:>12}: {stats['correct']:>3}/{stats['total']:>3} ({source_acc:>5.1f}%)") + + print(f"\n🔧 Configuration:") + print(f" Temperature: {temperature}") + print(f" Samples per problem: {n_sampling}") + print(f" Max tokens: {max_tokens}") + print(f" Top-p: {top_p}") + print(f" Seed: {seed}") + + print(f"\n📝 Token Statistics:") + print(f" Avg input tokens: {results['avg_input_tokens']:>10.1f}") + print(f" Avg output tokens: {results['avg_output_tokens']:>10.1f}") + print(f" Max input tokens: {results['max_input_tokens']:>10}") + print(f" Max output tokens: {results['max_output_tokens']:>10}") + + # Performance assessment for AIME + if accuracy >= 50: + tier = "🏆 EXCEPTIONAL" + elif accuracy >= 30: + tier = "✅ EXCELLENT" + elif accuracy >= 20: + tier = "🎯 VERY GOOD" + elif accuracy >= 10: + tier = "⚠️ GOOD" + elif accuracy >= 5: + tier = "📈 FAIR" + else: + tier = "❌ NEEDS IMPROVEMENT" + + print(f"\n🎖️ AIME Performance: {tier} ({accuracy:.1f}%)") + print(f"\n💾 Detailed results saved to: {filename}") + print(f"\n{'='*70}") + + return results + + +# Comparison functions for multiple model results +def compare_aime_results(all_results): + """Generate comprehensive comparison for AIME evaluation results""" + print(f"\n{'='*80}") + print("COMPREHENSIVE AIME MODEL COMPARISON") + print(f"{'='*80}") + + # Main comparison table + print(f"{'Model':<15} {'Accuracy %':<12} {'Pass@K %':<10} {'Correct':<8} {'Total':<8}") + print("-" * 80) + + for result in all_results: + print(f"{result['model_type']:<15} " + f"{result['accuracy']:<12.1f} " + f"{result['pass_at_k']:<10.1f} " + f"{result['correct_answers']:<8} " + f"{result['total_problems']:<8}") + + # Performance improvement analysis + if len(all_results) > 1: + print(f"\n{'='*50}") + print("IMPROVEMENT ANALYSIS") + print(f"{'='*50}") + + base_result = all_results[0] # Assume first is base model + + for i, result in enumerate(all_results[1:], 1): + print(f"\n{result['model_type']} vs {base_result['model_type']}:") + + accuracy_improvement = result['accuracy'] - base_result['accuracy'] + pass_k_improvement = result['pass_at_k'] - base_result['pass_at_k'] + + print(f" Accuracy improvement: {accuracy_improvement:+.1f}%") + print(f" Pass@K improvement: {pass_k_improvement:+.1f}%") + + # Dataset breakdown + print(f"\n{'='*50}") + print("PERFORMANCE BY DATASET") + print(f"{'='*50}") + + # Get all unique datasets from the first result + if all_results and 'source_accuracies' in all_results[0]: + datasets = list(all_results[0]['source_accuracies'].keys()) + + print(f"{'Model':<15}", end="") + for dataset in datasets: + print(f"{dataset:<15}", end="") + print() + print("-" * (15 + 15 * len(datasets))) + + for result in all_results: + print(f"{result['model_type']:<15}", end="") + for dataset in datasets: + accuracy = result['source_accuracies'].get(dataset, 0) + print(f"{accuracy:<15.1f}", end="") + print() + + # Save comparison + comparison_data = { + "summary": all_results, + "best_model": max(all_results, key=lambda x: x['accuracy']), + } + + with open("aime_model_comparison.json", "w") as f: + json.dump(comparison_data, f, indent=4) + + print(f"\nBest performing model: {comparison_data['best_model']['model_type']} " + f"({comparison_data['best_model']['accuracy']:.1f}% accuracy)") diff --git a/tests/utils/cleanup_utils.py b/tests/utils/cleanup_utils.py new file mode 100644 index 0000000000..e9163b1314 --- /dev/null +++ b/tests/utils/cleanup_utils.py @@ -0,0 +1,209 @@ +import gc +import logging +import os +import shutil +import torch +import sys +import warnings + +def clear_memory(variables_to_clear=None, verbose=False, clear_all_caches=True): + """ + Comprehensive memory clearing for persistent memory leaks. + + Args: + variables_to_clear: List of variable names to clear + verbose: Print memory status + clear_all_caches: Clear all types of caches (recommended for memory leaks) + """ + + # Save current logging levels + saved_log_levels = {} + for name, logger in logging.Logger.manager.loggerDict.items(): + if isinstance(logger, logging.Logger): + saved_log_levels[name] = logger.level + root_level = logging.getLogger().level + + if variables_to_clear is None: + variables_to_clear = ["inputs", "model", "base_model", "processor", "tokenizer", + "base_processor", "base_tokenizer", "trainer", + "peft_model", "bnb_config"] + + # 1. Clear LRU caches FIRST (very important for memory leaks) + if clear_all_caches: + clear_all_lru_caches(verbose) + + # 2. Delete specified variables + g = globals() + deleted_vars = [] + for var in variables_to_clear: + if var in g: + del g[var] + deleted_vars.append(var) + + if verbose and deleted_vars: + print(f"Deleted variables: {deleted_vars}") + + # 3. Multiple garbage collection passes (important for circular references) + for i in range(3): + collected = gc.collect() + if verbose and collected > 0: + print(f"GC pass {i+1}: collected {collected} objects") + + # 4. CUDA cleanup + if torch.cuda.is_available(): + # Get memory before cleanup + if verbose: + mem_before = torch.cuda.memory_allocated() / 1024**3 + + torch.cuda.empty_cache() + torch.cuda.synchronize() + + # Additional CUDA cleanup for persistent leaks + if clear_all_caches: + # Reset memory stats + torch.cuda.reset_peak_memory_stats() + torch.cuda.reset_accumulated_memory_stats() + + # Clear JIT cache + if hasattr(torch.jit, '_state') and hasattr(torch.jit._state, '_clear_class_state'): + torch.jit._state._clear_class_state() + + # Force another CUDA cache clear + torch.cuda.empty_cache() + + # Final garbage collection + gc.collect() + + if verbose: + mem_after = torch.cuda.memory_allocated() / 1024**3 + mem_reserved = torch.cuda.memory_reserved() / 1024**3 + print(f"GPU memory - Before: {mem_before:.2f} GB, After: {mem_after:.2f} GB") + print(f"GPU reserved memory: {mem_reserved:.2f} GB") + if mem_before > 0: + print(f"Memory freed: {mem_before - mem_after:.2f} GB") + + # restore original logging levels + logging.getLogger().setLevel(root_level) + for name, level in saved_log_levels.items(): + if name in logging.Logger.manager.loggerDict: + logger = logging.getLogger(name) + logger.setLevel(level) + +def clear_all_lru_caches(verbose=True): + """Clear all LRU caches in loaded modules.""" + cleared_caches = [] + + # Modules to skip to avoid warnings + skip_modules = { + 'torch.distributed', + 'torchaudio', + 'torch._C', + 'torch.distributed.reduce_op', + 'torchaudio.backend', + } + + # Create a static list of modules to avoid RuntimeError + modules = list(sys.modules.items()) + + # Method 1: Clear caches in all loaded modules + for module_name, module in modules: + if module is None: + continue + + # Skip problematic modules + if any(module_name.startswith(skip) for skip in skip_modules): + continue + + try: + # Look for functions with lru_cache + for attr_name in dir(module): + try: + # Suppress warnings when checking attributes + with warnings.catch_warnings(): + warnings.simplefilter("ignore", FutureWarning) + warnings.simplefilter("ignore", UserWarning) + warnings.simplefilter("ignore", DeprecationWarning) + + attr = getattr(module, attr_name) + if hasattr(attr, 'cache_clear'): + attr.cache_clear() + cleared_caches.append(f"{module_name}.{attr_name}") + except Exception: + continue # Skip problematic attributes + except Exception: + continue # Skip problematic modules + + # Method 2: Clear specific known caches + known_caches = [ + 'transformers.utils.hub.cached_file', + 'transformers.tokenization_utils_base.get_tokenizer', + 'torch._dynamo.utils.counters', + ] + + for cache_path in known_caches: + try: + parts = cache_path.split('.') + module = sys.modules.get(parts[0]) + if module: + obj = module + for part in parts[1:]: + obj = getattr(obj, part, None) + if obj is None: + break + if obj and hasattr(obj, 'cache_clear'): + obj.cache_clear() + cleared_caches.append(cache_path) + except Exception: + continue # Skip problematic caches + + if verbose and cleared_caches: + print(f"Cleared {len(cleared_caches)} LRU caches") + + +def clear_specific_lru_cache(func): + """Clear cache for a specific function.""" + if hasattr(func, 'cache_clear'): + func.cache_clear() + return True + return False + + +# Additional utility for monitoring cache sizes +def monitor_cache_sizes(): + """Monitor LRU cache sizes across modules.""" + cache_info = [] + + for module_name, module in sys.modules.items(): + if module is None: + continue + try: + for attr_name in dir(module): + try: + attr = getattr(module, attr_name) + if hasattr(attr, 'cache_info'): + info = attr.cache_info() + cache_info.append({ + 'function': f"{module_name}.{attr_name}", + 'size': info.currsize, + 'hits': info.hits, + 'misses': info.misses + }) + except: + pass + except: + pass + + return sorted(cache_info, key=lambda x: x['size'], reverse=True) + + +def safe_remove_directory(path): + try: + if os.path.exists(path) and os.path.isdir(path): + shutil.rmtree(path) + return True + else: + print(f"Path {path} is not a valid directory") + return False + except Exception as e: + print(f"Failed to remove directory {path}: {e}") + return False diff --git a/tests/utils/ocr_eval.md b/tests/utils/ocr_eval.md new file mode 100644 index 0000000000..e3e0619f6b --- /dev/null +++ b/tests/utils/ocr_eval.md @@ -0,0 +1,109 @@ + +# OCR Model Evaluator +A comprehensive Python module for evaluating Optical Character Recognition (OCR) models using Word Error Rate (WER) and Character Error Rate (CER) metrics. This evaluator supports vision-language models and provides detailed analysis with comparison capabilities across multiple models + +## Basic Usage + +```python +from ocr_evaluator import evaluate_ocr_model + +# Simple evaluation +avg_wer, avg_cer = evaluate_ocr_model( + model=your_model, + processor=your_processor, + dataset=your_dataset, + output_dir="evaluation_results" +) + +print(f"Average WER: {avg_wer:.4f}") +print(f"Average CER: {avg_cer:.4f}") +``` + + +### Dataset Format + +The evaluator expects datasets in a chatml conversational format with the following structure: +``` +dataset = [ + { + "messages": [ + { + "role": "system", + "content": [{"type": "text", "text": "You are an OCR system."}] + }, + { + "role": "user", + "content": [ + {"type": "text", "text": "Extract text from this image"}, + {"type": "image", "image": PIL_Image_object} + ] + }, + { + "role": "assistant", + "content": [{"type": "text", "text": "Ground truth text"}] + } + ] + }, + # ... more samples +] +``` + + +## Examples + +### Document OCR evaluation + +```python +from ocr_evaluator import OCRModelEvaluator +from datasets import load_dataset + +# Load document OCR dataset +dataset = load_dataset("your-ocr-dataset", split="test") + +# Convert to required format +eval_data = [format_document_sample(sample) for sample in dataset] + +# Evaluate models +evaluator = OCRModelEvaluator() + +# Compare different model configurations +configs = { + "Standard Model": {"temperature": 1.0, "max_new_tokens": 512}, + "Conservative Model": {"temperature": 0.7, "max_new_tokens": 256}, + "Creative Model": {"temperature": 1.5, "max_new_tokens": 1024} +} + +for config_name, params in configs.items(): + wer, cer = evaluator.evaluate_model( + model=base_model, + processor=processor, + dataset=eval_data, + output_dir=f"document_ocr_{config_name.lower().replace(' ', '_')}", + **params + ) + evaluator.add_to_comparison(config_name, wer, cer) + +# Generate final report +evaluator.print_model_comparison() +``` + +### Handwritting Recognition +```python +# Specialized evaluation for handwriting +def evaluate_handwriting_models(models, handwriting_dataset): + evaluator = OCRModelEvaluator() + + for model_name, (model, processor) in models.items(): + # Adjust parameters for handwriting recognition + wer, cer = evaluator.evaluate_model( + model=model, + processor=processor, + dataset=handwriting_dataset, + temperature=1.2, # Slightly higher for handwriting variety + max_new_tokens=128, # Usually shorter text + output_dir=f"handwriting_{model_name}" + ) + evaluator.add_to_comparison(f"Handwriting - {model_name}", wer, cer) + + return evaluator.print_model_comparison() +``` diff --git a/tests/utils/ocr_eval.py b/tests/utils/ocr_eval.py new file mode 100644 index 0000000000..5df476c1e3 --- /dev/null +++ b/tests/utils/ocr_eval.py @@ -0,0 +1,352 @@ +""" +OCR Model Evaluation Module + +This module provides functionality to evaluate OCR models on datasets with +word error rate (WER) and character error rate (CER) metrics. +""" + +import os +import torch +from tqdm import tqdm +import pandas as pd +from jiwer import wer, cer +from qwen_vl_utils import process_vision_info +import matplotlib.pyplot as plt +from typing import List, Dict, Tuple, Optional, Any +import traceback + + +class OCRModelEvaluator: + """ + A comprehensive OCR model evaluator that supports multiple models and provides + detailed analysis with WER and CER metrics. + """ + + def __init__(self): + """Initialize the OCR evaluator.""" + self.model_comparison_results = {} + + def evaluate_model( + self, + model: Any, + processor: Any, + dataset: List[Dict], + output_dir: str = "ocr_evaluation_results", + max_new_tokens: int = 1024, + temperature: float = 1.5, + min_p: float = 0.1, + verbose: bool = True + ) -> Tuple[Optional[float], Optional[float]]: + """ + Evaluate a model on an OCR dataset. + """ + # Create output directory if it doesn't exist + os.makedirs(output_dir, exist_ok=True) + + # Initialize results storage + results = [] + + # Process each sample in the dataset + for i, sample in enumerate(tqdm(dataset, desc="Evaluating OCR performance", disable=not verbose)): + try: + # Extract components from sample + messages = sample['messages'] + + # Get ground truth, image, and question + ground_truth, image, question, input_messages = self._extract_sample_components( + messages, i, verbose + ) + + if ground_truth is None or image is None or question is None: + continue + + # Generate model response + generated_response = self._generate_response( + model, processor, input_messages, max_new_tokens, temperature, min_p + ) + + # Calculate metrics + word_error = wer(ground_truth, generated_response) + char_error = cer(ground_truth, generated_response) + + # Save individual result + self._save_individual_result( + output_dir, i, question, generated_response, ground_truth, word_error, char_error + ) + + # Store results for summary + results.append({ + 'sample_id': i, + 'wer': word_error, + 'cer': char_error, + 'model_output': generated_response.strip(), + 'ground_truth': ground_truth, + 'question': question + }) + + except Exception as e: + if verbose: + print(f"Error processing sample {i}: {str(e)}") + traceback.print_exc() + + # Generate summary report + return self._generate_summary_report(results, output_dir, verbose) + + def _extract_sample_components( + self, + messages: List[Dict], + sample_idx: int, + verbose: bool + ) -> Tuple[Optional[str], Optional[Any], Optional[str], List[Dict]]: + """Extract ground truth, image, question, and input messages from sample.""" + + # Extract system message (if present) + system_message = next((msg for msg in messages if msg['role'] == 'system'), None) + + # Extract user message with the image and question + user_message = next((msg for msg in messages if msg['role'] == 'user'), None) + if not user_message: + if verbose: + print(f"Skipping sample {sample_idx}: No user message found") + return None, None, None, [] + + # Extract assistant message with ground truth + assistant_message = next((msg for msg in messages if msg['role'] == 'assistant'), None) + if not assistant_message: + if verbose: + print(f"Skipping sample {sample_idx}: No assistant message (ground truth) found") + return None, None, None, [] + + # Extract ground truth text + ground_truth = None + for content_item in assistant_message['content']: + if content_item['type'] == 'text': + ground_truth = content_item['text'] + break + + if not ground_truth: + if verbose: + print(f"Skipping sample {sample_idx}: No text found in assistant message") + return None, None, None, [] + + # Extract image and question from user message + image = None + question = None + + for content_item in user_message['content']: + if content_item['type'] == 'image': + image = content_item['image'] + elif content_item['type'] == 'text': + question = content_item['text'] + + if not image: + if verbose: + print(f"Skipping sample {sample_idx}: No image found in user message") + return None, None, None, [] + + if not question: + if verbose: + print(f"Skipping sample {sample_idx}: No question found in user message") + return None, None, None, [] + + # Construct messages for the model input (excluding assistant message) + input_messages = [] + if system_message: + input_messages.append(system_message) + input_messages.append(user_message) + + return ground_truth, image, question, input_messages + + def _generate_response( + self, + model: Any, + processor: Any, + input_messages: List[Dict], + max_new_tokens: int, + temperature: float, + min_p: float + ) -> str: + """Generate response from the model.""" + + # Preparation for inference using Qwen's specific processing + text = processor.apply_chat_template( + input_messages, tokenize=False, add_generation_prompt=True + ) + + # Process vision info (images/videos) from messages + image_inputs, video_inputs = process_vision_info(input_messages) + + # Create model inputs + inputs = processor( + text=[text], + images=image_inputs, + videos=video_inputs, + padding=True, + return_tensors="pt" + ) + inputs = inputs.to(model.device) + + # Generate response + with torch.no_grad(): + generated_ids = model.generate( + **inputs, + max_new_tokens=max_new_tokens, + temperature=temperature, + min_p=min_p, + use_cache=True + ) + + # Extract only the generated part (not the input) + generated_ids_trimmed = [ + out_ids[len(in_ids):] for in_ids, out_ids in zip(inputs.input_ids, generated_ids) + ] + + # Decode the generated text + generated_response = processor.batch_decode( + generated_ids_trimmed, + skip_special_tokens=True, + clean_up_tokenization_spaces=False + )[0] + + return generated_response + + def _save_individual_result( + self, + output_dir: str, + sample_idx: int, + question: str, + generated_response: str, + ground_truth: str, + word_error: float, + char_error: float + ): + """Save individual sample result to file.""" + output_file = os.path.join(output_dir, f"sample_{sample_idx}.txt") + with open(output_file, 'w', encoding='utf-8') as f: + f.write(f"Sample {sample_idx}\n") + f.write(f"Question: {question}\n\n") + f.write(f"Model output:\n{generated_response.strip()}\n\n") + f.write(f"Ground truth:\n{ground_truth}\n\n") + f.write(f"WER: {word_error:.4f}, CER: {char_error:.4f}") + + def _generate_summary_report( + self, + results: List[Dict], + output_dir: str, + verbose: bool + ) -> Tuple[Optional[float], Optional[float]]: + """Generate and save summary report.""" + if not results: + if verbose: + print("No results to summarize.") + return None, None + + df = pd.DataFrame(results) + + # Calculate overall averages + avg_wer = df['wer'].mean() + avg_cer = df['cer'].mean() + + # Save average metrics + with open(os.path.join(output_dir, "avg_metrics.txt"), 'w') as f: + f.write(f"Average WER: {avg_wer:.4f}\n") + f.write(f"Average CER: {avg_cer:.4f}\n") + + # Save detailed results + df.to_csv(os.path.join(output_dir, "detailed_results.csv"), index=False) + + if verbose: + print("\nResults Summary:") + print(f"Average WER: {avg_wer:.4f}") + print(f"Average CER: {avg_cer:.4f}") + print(f"\nDetailed results saved to {output_dir}/") + + return avg_wer, avg_cer + + def add_to_comparison(self, model_name: str, wer: float, cer: float): + """Add model results to the comparison tracker.""" + self.model_comparison_results[model_name] = { + "wer": wer, + "cer": cer + } + + def print_model_comparison(self, save_csv: bool = True, save_plot: bool = True) -> Optional[pd.DataFrame]: + """Print a comparison of all models evaluated so far.""" + if not self.model_comparison_results: + print("No model results available for comparison") + return None + + print("\n==== MODEL COMPARISON REPORT ====") + + # Create a comparison dataframe + comparison_df = pd.DataFrame({ + "Model": list(self.model_comparison_results.keys()), + "WER": [results["wer"] for results in self.model_comparison_results.values()], + "CER": [results["cer"] for results in self.model_comparison_results.values()] + }) + + # Sort by WER (best performance first) + comparison_df = comparison_df.sort_values("WER") + + # Display the comparison table + print("\nComparison Table (sorted by WER):") + print(comparison_df.to_string(index=False)) + + # Save the comparison table + if save_csv: + comparison_file = "model_comparison_results.csv" + comparison_df.to_csv(comparison_file, index=False) + print(f"\nComparison table saved to {comparison_file}") + + # Generate a bar chart visualization + if save_plot: + self._create_comparison_plot(comparison_df) + + return comparison_df + + def _create_comparison_plot(self, comparison_df: pd.DataFrame): + """Create and save comparison plot.""" + plt.figure(figsize=(12, 6)) + + # Plot WER + plt.subplot(1, 2, 1) + plt.bar(comparison_df["Model"], comparison_df["WER"], color='skyblue') + plt.title('Word Error Rate Comparison') + plt.ylabel('WER (lower is better)') + plt.ylim(bottom=0) + plt.xticks(rotation=45, ha='right') + + # Plot CER + plt.subplot(1, 2, 2) + plt.bar(comparison_df["Model"], comparison_df["CER"], color='lightgreen') + plt.title('Character Error Rate Comparison') + plt.ylabel('CER (lower is better)') + plt.ylim(bottom=0) + plt.xticks(rotation=45, ha='right') + + plt.tight_layout() + plt.savefig('ocr_model_comparison.png') + plt.show() + + print(f"\nVisualization saved to ocr_model_comparison.png") + + def get_comparison_results(self) -> Dict[str, Dict[str, float]]: + """Get the current comparison results.""" + return self.model_comparison_results.copy() + + def clear_comparison_results(self): + """Clear all comparison results.""" + self.model_comparison_results.clear() + + +def evaluate_ocr_model(model, processor, dataset, output_dir="ocr_evaluation_results", **kwargs): + """ + Convenience function that maintains backward compatibility with the original function. + """ + evaluator = OCRModelEvaluator() + return evaluator.evaluate_model(model, processor, dataset, output_dir, **kwargs) + + +def create_evaluator(): + """Create a new OCR evaluator instance.""" + return OCRModelEvaluator() diff --git a/tests/utils/perplexity_eval.md b/tests/utils/perplexity_eval.md new file mode 100644 index 0000000000..df86b558ed --- /dev/null +++ b/tests/utils/perplexity_eval.md @@ -0,0 +1,20 @@ +# Language Model Perplexity Evaluator + +A Python module for evaluating language models using perplexity metrics with sliding window approach for long sequences. This evaluator provides efficient computation of perplexity scores across datasets with model comparison capabilities. + +## Basic Usage + +```python +from perplexity_evaluator import ppl_model, add_to_comparison, print_model_comparison + +# Simple perplexity evaluation +dataset = {"text": ["Your text samples here...", "Another text sample..."]} +perplexity = ppl_model(model, tokenizer, dataset) + +print(f"Model Perplexity: {perplexity:.4f}") + +# Add to comparison tracker +add_to_comparison("My Model", perplexity) +print_model_comparison() +``` + diff --git a/tests/utils/perplexity_eval.py b/tests/utils/perplexity_eval.py new file mode 100644 index 0000000000..fa297540fd --- /dev/null +++ b/tests/utils/perplexity_eval.py @@ -0,0 +1,75 @@ +from tqdm import tqdm +import torch +import pandas as pd + +model_comparison_results = {} +#return the perplexity of the model on the dataset +#The perplexity is computed on each example, individually, with a sliding window for examples longer than 512 tokens. + + +def ppl_model(model, tokenizer, dataset): + nlls = [] + max_length = 2048 + stride = 512 + for s in tqdm(range(len(dataset['text']))): + encodings = tokenizer(dataset['text'][s], return_tensors="pt") + seq_len = encodings.input_ids.size(1) + prev_end_loc = 0 + for begin_loc in range(0, seq_len, stride): + end_loc = min(begin_loc + max_length, seq_len) + trg_len = end_loc - prev_end_loc + input_ids = encodings.input_ids[:, begin_loc:end_loc].to("cuda") + target_ids = input_ids.clone() + target_ids[:, :-trg_len] = -100 + # Create attention mask based on pad token id + pad_token_id = tokenizer.pad_token_id if tokenizer.pad_token_id is not None else 0 + attention_mask = (input_ids != pad_token_id).long() + with torch.no_grad(): + outputs = model(input_ids, labels=target_ids, attention_mask=attention_mask) + neg_log_likelihood = outputs.loss + nlls.append(neg_log_likelihood) + prev_end_loc = end_loc + if end_loc == seq_len: + break + ppl = torch.exp(torch.stack(nlls).mean()) + return ppl + + +#-------------------------------------------------------------------- + + +## ----------- Reporting helper function ----------- ## + +# Create a simple function to add results to the comparison +def add_to_comparison(model_name, ppl): + """Add model results to the comparison tracker""" + model_comparison_results[model_name] = { + "ppl": ppl + } + #return model_comparison_results + + + +# Create a function to print the comparison report whenever needed +def print_model_comparison(): + """Print a comparison of all models evaluated so far""" + if not model_comparison_results: + print("No model results available for comparison") + return + + print("\n==== MODEL COMPARISON REPORT ====") + + # Create a comparison dataframe + comparison_df = pd.DataFrame({ + "Model": list(model_comparison_results.keys()), + #"Perplexity": [results["ppl"] for results in model_comparison_results.values()], + "Perplexity": [ + # Convert tensors to CPU and then to float if needed + results["ppl"].cpu().item() if torch.is_tensor(results["ppl"]) else results["ppl"] + for results in model_comparison_results.values() + ], + }) + + # Display the comparison table + print("\nComparison Table:") + print(comparison_df.to_string(index=False)) diff --git a/unsloth/save.py b/unsloth/save.py index 107b553cd9..e61026318a 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -1625,9 +1625,9 @@ def create_ollama_modelfile(tokenizer, gguf_location): pass def create_ollama_model( - username: str, - model_name: str, - tag: str, + username: str, + model_name: str, + tag: str, modelfile_path: str ): try: @@ -1711,7 +1711,7 @@ def push_to_ollama( with open(f"Modelfile_{model_name}", "w") as f: f.write(model_file) f.close() - + create_ollama_model( username=username, model_name=model_name, @@ -2320,7 +2320,7 @@ def unsloth_generic_save( ) elif save_method == "merged_4bit_forced": save_method = "merged_4bit" - + merge_and_overwrite_lora( get_model_name, model = model, @@ -2524,8 +2524,8 @@ def patch_saving_functions(model, vision = False): if not vision: if hasattr(model, "config"): # Counteract tokenizers - model.push_to_hub_merged = types.MethodType(unsloth_push_to_hub_merged, model) - model.save_pretrained_merged = types.MethodType(unsloth_save_pretrained_merged, model) + model.push_to_hub_merged = types.MethodType(unsloth_generic_push_to_hub_merged, model) + model.save_pretrained_merged = types.MethodType(unsloth_generic_save_pretrained_merged, model) model.push_to_hub_gguf = types.MethodType(unsloth_push_to_hub_gguf, model) model.save_pretrained_gguf = types.MethodType(unsloth_save_pretrained_gguf, model) model.push_to_hub_ggml = types.MethodType(unsloth_convert_lora_to_ggml_and_push_to_hub, model)