diff --git a/README.md b/README.md index 95492b6ef3..4c4f8661bf 100644 --- a/README.md +++ b/README.md @@ -294,7 +294,6 @@ trainer = SFTTrainer( train_dataset = dataset, tokenizer = tokenizer, args = SFTConfig( - dataset_text_field = "text", max_seq_length = max_seq_length, per_device_train_batch_size = 2, gradient_accumulation_steps = 4, diff --git a/pyproject.toml b/pyproject.toml index 7fb572b735..523794ee1a 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -550,6 +550,48 @@ cu128-ampere-torch270 = [ "unsloth[flashattention]", ] +flashattentiontorch260abiFALSEcu12x = [ + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp39-cp39-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.9'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp310-cp310-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.10'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiFALSE-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13'", +] +flashattentiontorch260abiTRUEcu12x = [ + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiTRUE-cp39-cp39-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.9'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiTRUE-cp310-cp310-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.10'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiTRUE-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiTRUE-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.6cxx11abiTRUE-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13'", +] +flashattentiontorch250abiFALSEcu12x = [ + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiFALSE-cp39-cp39-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.9'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiFALSE-cp310-cp310-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.10'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiFALSE-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiFALSE-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiFALSE-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13'", +] +flashattentiontorch250abiTRUEcu12x = [ + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiTRUE-cp39-cp39-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.9'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiTRUE-cp310-cp310-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.10'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiTRUE-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiTRUE-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.5cxx11abiTRUE-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13'", +] +flashattentiontorch240abiFALSEcu12x = [ + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp39-cp39-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.9'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp310-cp310-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.10'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiFALSE-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13'", +] +flashattentiontorch240abiTRUEcu12x = [ + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp39-cp39-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.9'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp310-cp310-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.10'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp311-cp311-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.11'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp312-cp312-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.12'", + "flash-attn @ https://github.com/Dao-AILab/flash-attention/releases/download/v2.7.4.post1/flash_attn-2.7.4.post1+cu12torch2.4cxx11abiTRUE-cp313-cp313-linux_x86_64.whl ; platform_system == 'Linux' and python_version == '3.13'", +] intel-gpu-torch260 = [ "unsloth[huggingface]", diff --git a/unsloth/_auto_install.py b/unsloth/_auto_install.py index 308bf075e9..c8559394ed 100644 --- a/unsloth/_auto_install.py +++ b/unsloth/_auto_install.py @@ -18,7 +18,9 @@ from packaging.version import Version as V v = V(torch.__version__) cuda = str(torch.version.cuda) is_ampere = torch.cuda.get_device_capability()[0] >= 8 -if cuda != "12.1" and cuda != "11.8" and cuda != "12.4" and cuda != "12.6" and cuda != "12.8": raise RuntimeError(f"CUDA = {cuda} not supported!") +USE_ABI = torch._C._GLIBCXX_USE_CXX11_ABI +if cuda not in ("11.8", "12.1", "12.4", "12.6", "12.8"): + raise RuntimeError(f"CUDA = {cuda} not supported!") if v <= V('2.1.0'): raise RuntimeError(f"Torch = {v} too old!") elif v <= V('2.1.1'): x = 'cu{}{}-torch211' elif v <= V('2.1.2'): x = 'cu{}{}-torch212' diff --git a/unsloth/models/rl.py b/unsloth/models/rl.py index 5d2270810c..b385dba2eb 100644 --- a/unsloth/models/rl.py +++ b/unsloth/models/rl.py @@ -43,6 +43,7 @@ torch_compile_options = { "triton.cudagraphs" : False, } +from trl import __version__ as trl_version def vLLMSamplingParams(**kwargs): from vllm import SamplingParams @@ -440,6 +441,9 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): "torch_empty_cache_steps" : 250, "logging_steps" : 1, "max_seq_length" : None, + "num_generations" : 8, + "top_k" : None, + "vllm_mode" : "colocate", } for k, v in replacements.items(): x = f"{k}( = [^,\n]{{1,}})?,\n" @@ -542,7 +546,12 @@ def _patch_trl_rl_trainers(trainer_file = "grpo_trainer"): selective_log_softmax_code = selective_log_softmax_code, ) - + + if RLTrainer_name == "SFTTrainer": + original_text = 'self._signature_columns = ["input_ids", "attention_mask", "completion_mask"]' + new_text = 'self._signature_columns = ["input_ids", "attention_mask", "completion_mask","labels"]' + RLTrainer_source = RLTrainer_source.replace(original_text, new_text) + # Remove multiple doc strings if __RLConfig_doc__ != "" and RLTrainer_source.count(__RLTrainer_doc__) == 2: RLTrainer_source = RLTrainer_source.replace(__RLTrainer_doc__, "", 1) @@ -594,9 +603,15 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import if len(replacer) != 0: replacer = replacer[0] vllm_setter = "\n" + " "*8 + \ - "if hasattr(model, 'vllm_engine') and "\ - "hasattr(args, 'use_vllm') and (getattr(args, 'use_vllm', False) == False): "\ - "args.use_vllm = True\n" + "if hasattr(model, 'vllm_engine') and hasattr(args, 'use_vllm'):\n" + \ + " " * 12 + "if (getattr(args, 'use_vllm', False) == False):\n" + \ + " " * 16 + "args.use_vllm = True\n" + + if "grpo" in trainer_file and trl_version >= "0.18": + # If model has vllm_engine, then use vllm in colocate mode. Donot wait for server + vllm_setter += \ + " " * 12 + "args.vllm_mode='colocate'\n" + init = init.replace(replacer, replacer + vllm_setter) pass pass @@ -612,7 +627,8 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import if len(vllm_part) == 1: vllm_part, args = vllm_part[0][0], vllm_part[0][1] # Strip all comments - new_vllm_part = re.sub(r"\#[^\n]{1,}\n", "", vllm_part) + new_vllm_part = re.sub(r"^\s*\#[^\n]*\n?", "", vllm_part, flags=re.MULTILINE) # to also remove whole comment line instead of just starting at # + new_vllm_part = re.sub(r"\s*\#.*$", "", new_vllm_part, flags=re.MULTILINE) # remove comments that occur after code # Get SamplingParams sampling_params = re.findall( @@ -621,9 +637,9 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import new_vllm_part, flags = re.MULTILINE | re.DOTALL, ) + if len(sampling_params) == 1: sampling_params = sampling_params[0] - # Fix guided_decoding sampling_params = sampling_params.replace( "guided_decoding=guided_decoding,", @@ -635,11 +651,18 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import sampling_params = \ " "*12 + "self.llm = model.vllm_engine; self._last_loaded_step = 0; " + \ sampling_params # Add spaces + + # count the indentation of last line of sampling_params. + last_line = sampling_params.split("\n")[-1] + last_prev_line = sampling_params.split("\n")[-2] + last_prev_indentation = len(last_prev_line) - len(last_prev_line.lstrip()) + last_indentation = len(last_line) - len(last_line.lstrip()) + # Add extra arguments to SamplingParams extra = "**getattr(getattr(args, 'vllm_sampling_params', vLLMSamplingParams()), '_set_kwargs', {})" # Backwards replace - to_replace = "," + extra + "," + ")" + to_replace = ",\n" + " "*last_prev_indentation + extra + ",\n" + " "*last_indentation + ")" sampling_params = to_replace.join(sampling_params.rsplit(")", 1)) # Strip multiple commas sampling_params = re.sub(r"[\,][\s]{0,}\,", ",", sampling_params) @@ -647,9 +670,21 @@ def patch_functions(RLTrainer, trainer_file, RLTrainer_name, all_imports, import new_vllm_part = \ f"\n{' '*8}if {args}.use_vllm:\n{sampling_params}"\ f"\n{' '*8}else:\n" - - init = init.replace(vllm_part, new_vllm_part) pass + + if trl_version >= "0.18": + # Replace LLM init with already existing vLLM engine for colocate mode + vllm_llm_init_pattern = r"self\.llm\s*=\s*LLM\([^)]*\)*\)" + vllm_llm_replacement = "self.llm = model.vllm_engine\n" + new_vllm_part = re.sub( + vllm_llm_init_pattern, + vllm_llm_replacement, + new_vllm_part, + flags=re.DOTALL # Ensure . matches newlines [[5]] + ) + + init = init.replace(vllm_part, new_vllm_part) + pass # Search for vLLM calling in all child functions diff --git a/unsloth/models/rl_replacements.py b/unsloth/models/rl_replacements.py index 376d1e9a28..af401a2696 100644 --- a/unsloth/models/rl_replacements.py +++ b/unsloth/models/rl_replacements.py @@ -20,6 +20,7 @@ __all__ = [ "RL_METRICS_CHANGES", ] +import os import re import torch import inspect @@ -207,24 +208,34 @@ RL_FUNCTIONS["grpo_trainer"].append(grpo_trainer__move_model_to_vllm) def grpo_trainer__get_per_token_logps(function_name, function): if function_name != "_get_per_token_logps": return function - def _get_per_token_logps(self, model, input_ids, attention_mask, logits_to_keep): - if os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0': + def _get_per_token_logps(self, model, input_ids, attention_mask, logits_to_keep, calc_logprob_flag = None): + if os.environ.get('UNSLOTH_USE_NEW_MODEL', '0') == '0' and not calc_logprob_flag: return None # Unsloth efficient GRPO # Otherwise, calculate normally: if not hasattr(self, '_autocast_dtype'): self._autocast_dtype = torch.float16 if os.environ.get('ACCELERATE_MIXED_PRECISION', 'fp16') == 'fp16' else torch.bfloat16 if os.environ.get('UNSLOTH_FORCE_FLOAT32', '0') == '1': self._autocast_dtype = torch.float16 + + os.environ["UNSLOTH_RETURN_HIDDEN_STATES"] = "1" with torch.amp.autocast(device_type = 'cuda', dtype = self._autocast_dtype): # We add 1 to `logits_to_keep` because the last logits of the sequence is later excluded - logits = model(input_ids=input_ids, attention_mask=attention_mask, logits_to_keep=logits_to_keep + 1).logits - logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred - - input_ids = input_ids[:, -logits_to_keep:] + hidden_states = model(input_ids=input_ids, attention_mask=attention_mask, logits_to_keep=logits_to_keep + 1).logits + #logits = logits[:, :-1, :] # (B, L-1, V), exclude the last logit: it corresponds to the next token pred + return hidden_states + # input_ids = input_ids[:, -logits_to_keep:] # For transformers<=4.48, logits_to_keep argument isn't supported, so here we drop logits ourselves. # See https://github.com/huggingface/trl/issues/2770 - logits = logits[:, -logits_to_keep:] - return logits - # return selective_log_softmax(logits, input_ids) # compute logprobs for the input tokens + # logits = logits[:, -logits_to_keep:] + # return logits + # logps = selective_log_softmax(logits, input_ids) + + # row_indices, col_indices = torch.where(logps < -20) + + # # Method 1: Check if tensors have elements + # if len(row_indices) > 0 and len(col_indices) > 0: + # breakpoint() # Breakpoint triggered here + # print("Found high values!") + # return logps # compute logprobs for the input tokens pass pass @@ -264,7 +275,13 @@ def grpo_trainer_compute_loss(function_name, function): per_token_logps = self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep) # Compute the KL divergence between the model and the reference model - ref_per_token_logps = inputs["ref_per_token_logps"] + # _prepare_inputs doesn't return reference log probs anymore. We need to calculate it ourselves. + # https://github.com/huggingface/trl/blob/05bc43e960396581e458195b8388efe6b82cae1f/trl/trainer/grpo_trainer.py#L1328 + if self.beta != 0.0: + with torch.inference_mode(), model.disable_adapter(): + ref_per_token_logps = self._get_per_token_logps(model, input_ids, attention_mask, logits_to_keep) + else: + ref_per_token_logps = None # per_token_kl = torch.exp(ref_per_token_logps - per_token_logps) - (ref_per_token_logps - per_token_logps) - 1 # x - x.detach() allows for preserving gradients from x @@ -272,16 +289,35 @@ def grpo_trainer_compute_loss(function_name, function): # per_token_loss = torch.exp(per_token_logps - per_token_logps.detach()) * advantages.unsqueeze(1) # per_token_loss = -(per_token_loss - self.beta * per_token_kl) # loss = ((per_token_loss * completion_mask).sum(dim=1) / completion_mask.sum(dim=1)).mean() + if "old_per_token_logps" in inputs.keys(): + old_hidden_states = inputs["old_per_token_logps"] + else: + old_hidden_states = None input_ids = input_ids[:, -logits_to_keep:] if per_token_logps is not None: loss, completion_length, mean_kl = grpo_compute_loss_slow( - ref_per_token_logps, per_token_logps, input_ids, completion_mask, self.beta, advantages, + ref_per_token_logps, per_token_logps, old_hidden_states, input_ids, completion_mask, self.beta, advantages, + loss_type = self.args.loss_type, + epsilon_low = self.epsilon_low, epsilon_high = self.epsilon_high, + max_completion_length = self.args.max_completion_length, + delta = self.args.delta, ) else: - loss, completion_length, mean_kl = grpo_accumulated_loss( - self, _input_ids, logits_to_keep, completion_mask, advantages, - n_chunks = self.args.unsloth_num_chunks, - ) + if hasattr(self.args, "loss_type"): + loss, completion_length, mean_kl = grpo_accumulated_loss( + self, _input_ids, logits_to_keep, completion_mask, advantages, old_hidden_states, + n_chunks = self.args.unsloth_num_chunks, + loss_type = self.args.loss_type, + epsilon_low = self.epsilon_low, epsilon_high = self.epsilon_high, + max_completion_length = self.args.max_completion_length, + delta = self.args.delta, + ) + else: + # to ensure backwards compatibility with trl 0.15.2 and maybe even 0.17 + loss, completion_length, mean_kl = grpo_accumulated_loss( + self, _input_ids, logits_to_keep, completion_mask, advantages, old_hidden_states, + n_chunks = self.args.unsloth_num_chunks, + ) # Log the metrics # completion_length = self.accelerator.gather_for_metrics(completion_mask.sum(1)).float().mean().item() diff --git a/unsloth/save.py b/unsloth/save.py index 5652e15098..107b553cd9 100644 --- a/unsloth/save.py +++ b/unsloth/save.py @@ -888,11 +888,14 @@ def install_llama_cpp_old(version = -10): os.path.exists("llama.cpp/llama-quantize.exe") or os.path.exists("llama.cpp/llama-quantize") or os.path.exists("llama.cpp/quantize.exe") or - os.path.exists("llama.cpp/quantize") + os.path.exists("llama.cpp/quantize") or + os.path.exists("llama.cpp/build/bin/llama-quantize") or + os.path.exists("llama.cpp/build/bin/quantize") ): raise RuntimeError( "Unsloth: The file 'llama.cpp/llama-quantize' or `llama.cpp/quantize` does not exist.\n"\ - "But we expect this file to exist! Maybe the llama.cpp developers changed the name or check extension of the llama-quantize file." + "We've also double checked the building directory under 'llama.cpp/build/bin/'.\n"\ + "But we expect this file to exist! Check if the file exists under llama.cpp and investigate the building process of llama.cpp (make/cmake)!" ) pass pass @@ -1082,10 +1085,15 @@ def save_to_gguf( quantize_location = "llama.cpp/llama-quantize.exe" elif os.path.exists("llama.cpp/llama-quantize"): quantize_location = "llama.cpp/llama-quantize" + elif os.path.exists("llama.cpp/build/bin/llama-quantize"): + quantize_location = "llama.cpp/build/bin/llama-quantize" + elif os.path.exists("llama.cpp/build/bin/quantize"): + quantize_location = "llama.cpp/build/bin/quantize" else: raise RuntimeError( - "Unsloth: The file ('llama.cpp/llama-quantize' or 'llama.cpp/llama-quantize.exe' if you are on Windows WSL) or 'llama.cpp/quantize' does not exist.\n"\ - "But we expect this file to exist! Maybe the llama.cpp developers changed the name or check extension of the llama-quantize file." + "Unsloth: The file 'llama.cpp/llama-quantize' or `llama.cpp/quantize` does not exist.\n"\ + "We've also double checked the building directory under 'llama.cpp/build/bin/'.\n"\ + "But we expect this file to exist! Check if the file exists under llama.cpp and investigate the building process of llama.cpp (make/cmake)!" ) pass diff --git a/unsloth/trainer.py b/unsloth/trainer.py index 012be4b0cb..75fdd410e1 100644 --- a/unsloth/trainer.py +++ b/unsloth/trainer.py @@ -194,8 +194,15 @@ def _backwards_compatible_trainer(trainer_class, config_class): config_dict.update(additional_config_kwargs) # Create Config with all the collected parameters - config = config_class(**config_dict) - + # Reinitialising config class with parameters (that were none initially but populated on first init) + # causes the 2nd init to fail as there are mutual exclusive checks on pairs of parameters. + # Refer: https://github.com/huggingface/trl/blob/main/trl/trainer/grpo_config.py#L499-L502 for example + # So we only create config class if the previous init was not TrainingArguments + if not isinstance(training_args, TrainingArguments): + config = config_class(**config_dict) + else: + config = training_args + # Reconstruct kwargs for Trainer kwargs = trainer_kwargs kwargs["args"] = config