* Update rl.py * Fix CE Loss * Versioning * Update loader.py * Update loader.py * extract_model_type_from_config * Model types * Update loader.py * get_transformers_model_type * Update loader.py * Update loader.py * Update loader.py * Update rl.py * Update pyproject.toml * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Versioning * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Update vision.py * Update vision.py * Fix DataParallel * Update _utils.py * Update rl.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update mapper.py * Versioning * Update loader.py * Update loader.py * Update rl.py * Versioning * Update _utils.py * Fix auto_mapping * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update loader.py * Message * Update vision.py * Update loader.py * Update vision.py * cache_implementation * Update vision.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Save max_seq_length * Update _utils.py * Update rl.py * Update vision.py * Update llama.py * Mistral3 vllm (#3349) * [WIP] use vLLM for vision language models * Update README.md Editing icon sizes * Update README.md Updating icon sizes * Update README.md (#2885) * MoE kernels AGPLv3 * versioning * Many bug fixes (#2908) * add deepseek v3 * add deepseek r1 base * add deepseek r1 zero * add deepseek distill llama * add deepseek distill models * remove redundant code when constructing model names * add mistral small to registry * rename model registration methods * rename deepseek registration methods * refactor naming for mistral and phi * add global register models * refactor model registration tests for new registry apis * add model search method * remove deprecated registration api * add quant type test * add registry readme * make llama registration more specific * clear registry when executing individual model registration file * more registry readme updates * Update _auto_install.py * Llama4 * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Synthetic data * Update mapper.py * Xet and Synthetic * Update synthetic.py * Update loader.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update _utils.py * Update pyproject.toml * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update chat_templates.py * Seasame force float16 / float32 * Fix Seasame * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * is_multimodal * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * UNSLOTH_DISABLE_STATIC_GENERATION * Update vision.py * Auto vision detection * Sesame * Whisper * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update _utils.py * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * logging * Update pyproject.toml * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * logits / temperature * Update rl_replacements.py * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Debugging only * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Generic efficient GRPO * Update rl_replacements.py * Update rl_replacements.py * Remove debugging * Update rl_replacements.py * Update rl_replacements.py * Update vision.py * Update llama.py * Update rl_replacements.py * versioning * Update _utils.py * Update vision.py * Update mapper.py * Update loader.py * Update mapper.py * Update vision.py * Update loader.py * Update vision.py * Update loader.py * Update _utils.py * Update vision.py * gradient checkpointing * Gemma 3N fixes * Update loader.py * Versioning * Gemma 3N fixes * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Fix setup.py * setup.py * Prints * Update setup.py * Update setup.py * Update setup.py * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update vision.py * Update vision.py * Update pyproject.toml * Update vision.py * Update _utils.py * Update __init__.py * Update __init__.py --------- Co-authored-by: jeromeku <jerome.ku@gmail.com> Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> * silienty skip falcon h1 import is transformers_version < 4.53.0 (#2912) * Dynamically adjust get_per_token_logps function and patch as well (#2911) * add intel gpu with vllm support (#2903) * [bugs] fix for casual mask (#2868) * fix for casual mask * use un_casual in sdpa * add missing mask * fix for type * Explicitly check if xformers exists for attention (#2889) * Update __init__.py * Update llama.py * if mlp doesn't exist in layer module check for feed_forward name for falcon h1 (#2913) * Move inputs to right devices. (#2919) * Move tensors to right devices * fix multi gpu for non mistral models * multi GPU RoPE for gemma2 * Finish up multi GPU inference * Make multiGPU rope a list * Remove unnecessary transfer to CPU * Remove unnecessary move to CPU * Donot move inputs to device yet will be handled separately in another PR * Move inputs to appropriate decoder device * Make device count global variable * Cleanup RoPE device code * Fixup num_gpu to device count * Cleanup device counts * Use device index for RoPE get_cache * Donot typecast * Use tuple instead of list for tensors. Use device index directly * fixup move to device logic * WIP VLM vLLM * Make vLLM patch a function * Add save and load lora functions * Make fast_inference setup depend on the flag * Improve fast inference patching mechanism * Make vision setting depend on checks in fastbasemodel * Check LoRA and vLLM intercompatibility for vision models * Comment pointing to vLLM LoRA check * Improve lora validation on vLLM * Error out on no vLLM and increase max lora rank * Bug fixes (#3017) * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update _utils.py * Update pyproject.toml * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update chat_templates.py * Seasame force float16 / float32 * Fix Seasame * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * is_multimodal * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * UNSLOTH_DISABLE_STATIC_GENERATION * Update vision.py * Auto vision detection * Sesame * Whisper * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update _utils.py * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * logging * Update pyproject.toml * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * logits / temperature * Update rl_replacements.py * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Debugging only * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Generic efficient GRPO * Update rl_replacements.py * Update rl_replacements.py * Remove debugging * Update rl_replacements.py * Update rl_replacements.py * Update vision.py * Update llama.py * Update rl_replacements.py * versioning * Update _utils.py * Update vision.py * Update mapper.py * Update loader.py * Update mapper.py * Update vision.py * Update loader.py * Update vision.py * Update loader.py * Update _utils.py * Update vision.py * gradient checkpointing * Gemma 3N fixes * Update loader.py * Versioning * Gemma 3N fixes * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Fix setup.py * setup.py * Prints * Update setup.py * Update setup.py * Update setup.py * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update vision.py * Update vision.py * Update pyproject.toml * Update vision.py * Update _utils.py * Update __init__.py * Update __init__.py * Small fixes * Update vision.py * Update vision.py * versioning * Update __init__.py * Update llama.py * Update rl.py * Update rl.py * Update _utils.py * Update vision.py * Update vision.py * compiler stance * Update _utils.py * Update pyproject.toml * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Revert "Revert "Add Qwen2.5-VL-32B-Instruct mapping to fix quantized model me…" (#2990) This reverts commit6631007493. * skip_guard_eval_unsafe fix * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update llama.py * Update llama.py * Fix `quantization_method` * versioning * fix for casual mask (#3011) * [intel] add for intel path for llama.py (#3012) * fix for intel path * remove unuse code * Update unsloth/models/llama.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * Update llama.py * Fix Gemma 2 (#3024) * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update _utils.py * Update pyproject.toml * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update chat_templates.py * Seasame force float16 / float32 * Fix Seasame * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * is_multimodal * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * UNSLOTH_DISABLE_STATIC_GENERATION * Update vision.py * Auto vision detection * Sesame * Whisper * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update _utils.py * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * logging * Update pyproject.toml * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * logits / temperature * Update rl_replacements.py * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Debugging only * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Generic efficient GRPO * Update rl_replacements.py * Update rl_replacements.py * Remove debugging * Update rl_replacements.py * Update rl_replacements.py * Update vision.py * Update llama.py * Update rl_replacements.py * versioning * Update _utils.py * Update vision.py * Update mapper.py * Update loader.py * Update mapper.py * Update vision.py * Update loader.py * Update vision.py * Update loader.py * Update _utils.py * Update vision.py * gradient checkpointing * Gemma 3N fixes * Update loader.py * Versioning * Gemma 3N fixes * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Fix setup.py * setup.py * Prints * Update setup.py * Update setup.py * Update setup.py * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update vision.py * Update vision.py * Update pyproject.toml * Update vision.py * Update _utils.py * Update __init__.py * Update __init__.py * Small fixes * Update vision.py * Update vision.py * versioning * Update __init__.py * Update llama.py * Update rl.py * Update rl.py * Update _utils.py * Update vision.py * Update vision.py * compiler stance * Update _utils.py * Update pyproject.toml * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Revert "Revert "Add Qwen2.5-VL-32B-Instruct mapping to fix quantized model me…" (#2990) This reverts commit6631007493. * skip_guard_eval_unsafe fix * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update llama.py * Update llama.py * Fix `quantization_method` * versioning * Update _utils.py * Update _utils.py * Update _utils.py * falcon force float32 on sm<75 machines (#3026) * Fix torch compile issues (#3028) * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update pyproject.toml * Delete .gitignore * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update _utils.py * Update pyproject.toml * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update chat_templates.py * Seasame force float16 / float32 * Fix Seasame * Update loader.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * is_multimodal * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update vision.py * Update vision.py * Update vision.py * UNSLOTH_DISABLE_STATIC_GENERATION * Update vision.py * Auto vision detection * Sesame * Whisper * Update loader.py * Update loader.py * Update loader.py * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update loader.py * Update loader.py * Update loader.py * Update loader.py * Update _utils.py * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * logging * Update pyproject.toml * Update rl.py * versioning * Update rl.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * logits / temperature * Update rl_replacements.py * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Debugging only * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Generic efficient GRPO * Update rl_replacements.py * Update rl_replacements.py * Remove debugging * Update rl_replacements.py * Update rl_replacements.py * Update vision.py * Update llama.py * Update rl_replacements.py * versioning * Update _utils.py * Update vision.py * Update mapper.py * Update loader.py * Update mapper.py * Update vision.py * Update loader.py * Update vision.py * Update loader.py * Update _utils.py * Update vision.py * gradient checkpointing * Gemma 3N fixes * Update loader.py * Versioning * Gemma 3N fixes * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Fix setup.py * setup.py * Prints * Update setup.py * Update setup.py * Update setup.py * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update pyproject.toml * Update vision.py * Update vision.py * Update pyproject.toml * Update vision.py * Update _utils.py * Update __init__.py * Update __init__.py * Small fixes * Update vision.py * Update vision.py * versioning * Update __init__.py * Update llama.py * Update rl.py * Update rl.py * Update _utils.py * Update vision.py * Update vision.py * compiler stance * Update _utils.py * Update pyproject.toml * Update pyproject.toml * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Revert "Revert "Add Qwen2.5-VL-32B-Instruct mapping to fix quantized model me…" (#2990) This reverts commit6631007493. * skip_guard_eval_unsafe fix * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update synthetic.py * Update llama.py * Update llama.py * Fix `quantization_method` * versioning * Update _utils.py * Update _utils.py * Update _utils.py * check stride * Cleanup * Update rope_embedding.py * Update gemma2.py * Fix `set_stance` * Update pyproject.toml * Update _utils.py * Fixup patch vllm * Disable mllama * Use variables to decide VLM support * Better attn_impl handling * Patch TF protobuf incompatability * Torch 2.8 (#3186) * Fix mamba * Update loader.py * Update vision.py * Update loader.py * Filter vLLM standby logs (#3131) * filter vLLM standby logs * safeguard standby logger patch * Update unsloth/models/_utils.py * Update unsloth/models/_utils.py * Update unsloth/models/_utils.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * Update loader.py * Add scaler * Update llama.py * Update _utils.py * Versioning * GPT OSS fix * GPT OSS fix * Update loader.py * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Update vision.py * Update llama.py * Update llama.py * Update llama.py * Versioning * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Upcast norms * Update loader.py * Update vision.py * Upcast layernorms * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update save.py * Update rl.py * Update pyproject.toml * Update rl.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update _utils.py * Update __init__.py * Torch 2.8 * Update rl_replacements.py --------- Co-authored-by: Datta Nimmaturi <venkatadattasainimmaturi@gmail.com> * Update _auto_install.py * Update pyproject.toml * Update rl.py * Protobuf issue * Update pyproject.toml * Fix extras transformers typo in pyproject.toml * Update _utils.py * Bug fixes (#3195) * Fix mamba * Update loader.py * Update vision.py * Update loader.py * Filter vLLM standby logs (#3131) * filter vLLM standby logs * safeguard standby logger patch * Update unsloth/models/_utils.py * Update unsloth/models/_utils.py * Update unsloth/models/_utils.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * Update loader.py * Add scaler * Update llama.py * Update _utils.py * Versioning * GPT OSS fix * GPT OSS fix * Update loader.py * Update vision.py * Update vision.py * Update loader.py * Update vision.py * Update vision.py * Update llama.py * Update llama.py * Update llama.py * Versioning * Update mapper.py * Update vision.py * Update vision.py * Update vision.py * Upcast norms * Update loader.py * Update vision.py * Upcast layernorms * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update save.py * Update rl.py * Update pyproject.toml * Update rl.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update _utils.py * Update __init__.py * Torch 2.8 * Update rl_replacements.py * Update loader.py * UNSLOTH_ENABLE_CCE * Fix * Update loader.py * Update loader.py * Update __init__.py * Update __init__.py * Update __init__.py * Update __init__.py * Import fixes * Update loader.py * Fix aimv2 issue * Update loader.py * Update import_fixes.py * Update import_fixes.py * Update loader.py * Update loader.py * Update loader.py * Upgrade * Update loader.py * Update loader.py * Update loader.py * Update loader.py --------- Co-authored-by: Datta Nimmaturi <venkatadattasainimmaturi@gmail.com> * adallow float32 dtype in FastLanguageModel (#3204) * Update loader.py * Update vision.py * Suppress message and use unsloth sampling params * Use trl sampling params for now * Improve error message * fixup quantized fast inference model name * Add mistral 3 support --------- Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Co-authored-by: Daniel Han <danielhanchen@gmail.com> Co-authored-by: jeromeku <jerome.ku@gmail.com> Co-authored-by: DoubleMathew <mmathew23@gmail.com> Co-authored-by: Lei Zhenyuan <zhenyuan.lei@intel.com> Co-authored-by: parth2510 <parthguptapg7326@gmail.com> * Set padding to 0 * Fix patch * fixup patch (#3359) Co-authored-by: Datta Nimmaturi <venkatadattasainimmaturi@gmail.com> * Update vision.py * Versioning * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * MXFP4 dequant * Update loader.py * Update vision.py * load_in_16bit * Update vision.py * Update vision.py * Update vision.py * Update rl.py * Update vision.py * offload_embedding * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update rl_replacements.py * Update loader.py * Fix padding issue * Update pyproject.toml * Update _utils.py * Update pyproject.toml * Update _utils.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * Update vision.py * New models * Update llama.py * Versioning * Update _utils.py * Update llama.py * Update _utils.py * Update llama.py * Fix AMD * Update _utils.py * Update llama.py * Update vision.py * DEVICE_TYPE_TORCH * Update __init__.py * Update __init__.py * Update _utils.py * Move DEVICE_TYPE * Update rl_replacements.py * Update loader.py * AMD install script * Move AMD * Update _amd_install.sh * Update pyproject.toml * Update pyproject.toml * Delete _amd_install.sh * Update device_type.py * Update loader.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Update _utils.py * Update tokenizer_utils.py * Versioning * Update pyproject.toml * Update loader.py * Update _utils.py * Update pyproject.toml * Update pyproject.toml * Update _utils.py * Update pyproject.toml * Update _utils.py * Update _utils.py * Update loader.py * Update _utils.py * Update _utils.py * local_files_only * Cut Cross Entropy * Update llama.py * Update vision.py * Update vision.py * Update vision.py * Qwen 3 VL vLLM (#3489) * Update __init__.py * patch_torchao * torchao_logger * Update rl_replacements.py * Fix * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update _utils.py * Versioning * fbgemm fp8 block quant support (>=1.4.0) (#3531) * fbgemm fp8 block quant support (>=1.4.0) * Verify for fp8 support before proceeding * Use unsloth zoo's Version and improve comments * spacessss * Update vision.py * Update vision.py * Update rl.py * vllm_sampling_params * Update rl.py * Update rl.py * Update rl.py * Add `ruff` pre-commit hook and apply it (#3424) * Add Ruff pre-commit config and workflow * Add kwarg spacing enforcement helper * Apply Ruff formatting * Update fp8.py * Revert ruff on some files * Update * force-exclude = true * Datasets issue * Ruff * Remove mapper * Update mapper.py * Update pyproject.toml --------- Co-authored-by: Datta Nimmaturi <venkatadattasainimmaturi@gmail.com> Co-authored-by: Michael Han <107991372+shimmyshimmer@users.noreply.github.com> Co-authored-by: jeromeku <jerome.ku@gmail.com> Co-authored-by: DoubleMathew <mmathew23@gmail.com> Co-authored-by: Lei Zhenyuan <zhenyuan.lei@intel.com> Co-authored-by: parth2510 <parthguptapg7326@gmail.com> Co-authored-by: Dan Saunders <danjsaund@gmail.com>
374 lines
12 KiB
Python
374 lines
12 KiB
Python
"""
|
|
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()
|