From 4b7517978307ecfd93ec919a63f15fa371b20652 Mon Sep 17 00:00:00 2001 From: Daniel Han Date: Wed, 5 Mar 2025 00:51:10 -0800 Subject: [PATCH] SFT dataset prepare --- unsloth/models/rl_replacements.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 5ea61cb9b3..584214e801 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -78,6 +78,20 @@ def sft_trainer_prepare_dataset(function_name, function): if function_name != "_prepare_non_packed_dataloader" and \ function_name != "_prepare_dataset": return function + fast_sft_prepare_dataset = RL_REPLACEMENTS.get("sft_prepare_dataset", None) + if fast_sft_prepare_dataset is not None: + params = inspect.signature(fast_sft_prepare_dataset).parameters.keys() + params = ".*?".join(params) + matched = re.match( + r"[\s]{0,}def _prepare_dataset\(.*?" + params + r".*?\)", + function, + flags = re.MULTILINE | re.DOTALL, + ) + if matched: + # Use fast version! + return inspect.getsource(fast_sft_prepare_dataset) + pass + check_text = \ "if 'tokenizer' not in locals(): tokenizer = processing_class\n"\ "if 'formatting_func' not in locals(): raise RuntimeError('Unsloth: Please file a bug report - `formatting_func` does not exist!')\n"\