Fix bugs (#1706)
* Bug fixes * fix: flash_attn_detection_error (#1556) * fix: flash_attn_detection_error * Update _utils.py --------- Co-authored-by: Daniel Han <danielhanchen@gmail.com> * Update mapper.py * Update gemma.py * Update gemma.py * Update gemma.py * Update gemma.py * dim fix * Update _utils.py * Torch 2.6 support * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Faster inference? * Update llama.py * Update llama.py * Update utils.py * Update llama.py * Update llama.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update utils.py * Update mapper.py * Fast Inference via vLLM * Update llama.py * Update llama.py * Update utils.py * Create rl.py * PatchRL * Update rl.py * Update rl.py * Update rl.py * PatchRLStatistics * Update rl.py * Update rl.py * Update rl.py * Update utils.py * Update utils.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * RL metrics * Update rl.py * RL metrics * Update __init__.py * Update rl.py * Update rl.py * Update rl.py * Update chat_templates.py * Update mapper.py * Fp8 cache * Update llama.py * Update llama.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update __init__.py * Update loader.py * Update rl.py * Update rl.py * Update _utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Better TRL handling * Update rl.py * Update tokenizer_utils.py * Auto patching * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update rl.py * Update tokenizer_utils.py * Update rl.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update tokenizer_utils.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update tokenizer_utils.py * Update rl.py * Update rl.py * Update rl.py * max seq length * Update rl.py * Update rl.py * Patching * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * NEFTune * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Extra replacements * Update rl_replacements.py * Update rl.py * extra RL replacements * 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 llama.py * Update rl_replacements.py * Update _utils.py * Update loader_utils.py * Update rl.py * Update rl_replacements.py * Update rl_replacements.py * Update rl.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * autocast * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update pyproject.toml * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update rl_replacements.py * Update llama.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 * Update rl_replacements.py * Update llama.py * Update _utils.py * Update llama.py * Update _utils.py * Update rl_replacements.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update rl.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py * Update rl_replacements.py * Update llama.py * Update llama.py * Update llama.py * Update llama.py --------- Co-authored-by: Zhe Zhang <2631992879@qq.com>
This commit is contained in:
parent
aaa29ff83f
commit
c51f8b41db
4 changed files with 19 additions and 21 deletions
|
|
@ -12,7 +12,7 @@
|
|||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
__version__ = "2025.2.8"
|
||||
__version__ = "2025.2.9"
|
||||
|
||||
__all__ = [
|
||||
"SUPPORTS_BFLOAT16",
|
||||
|
|
|
|||
|
|
@ -708,7 +708,7 @@ def LlamaModel_fast_forward(
|
|||
if attention_mask is None:
|
||||
padding_mask = None
|
||||
elif self.training:
|
||||
# elif attention_mask is not None and self.training:
|
||||
# elif attention_mask is None:
|
||||
attention_mask = None
|
||||
padding_mask = None
|
||||
else:
|
||||
|
|
@ -724,7 +724,8 @@ def LlamaModel_fast_forward(
|
|||
past_key_values_length,
|
||||
sliding_window = getattr(self.config, "sliding_window", None),
|
||||
)
|
||||
attention_mask = attention_mask.to(torch.bool)
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.to(torch.bool)
|
||||
pass
|
||||
|
||||
hidden_states = inputs_embeds
|
||||
|
|
|
|||
|
|
@ -565,8 +565,8 @@ pass
|
|||
|
||||
|
||||
def PatchFastRL(algorithm = None, FastLanguageModel = None):
|
||||
return
|
||||
# if FastLanguageModel is not None: PatchRL(FastLanguageModel)
|
||||
# patch_trl_rl_trainers()
|
||||
# if algorithm is not None: PatchRLStatistics(algorithm)
|
||||
if FastLanguageModel is not None: PatchRL(FastLanguageModel)
|
||||
patch_trl_rl_trainers()
|
||||
if type(algorithm) is str and algorithm.islower():
|
||||
PatchRLStatistics(algorithm)
|
||||
pass
|
||||
|
|
|
|||
|
|
@ -101,23 +101,20 @@ RL_FUNCTIONS["sft_trainer"].append(sft_trainer_prepare_dataset)
|
|||
|
||||
# Ignore mean_token_accuracy since it needs logits
|
||||
# We override it directly with our version
|
||||
def _sft_trainer_compute_loss(self, model, inputs, return_outputs = False, num_items_in_batch = None):
|
||||
(loss, outputs) = super().compute_loss(
|
||||
model,
|
||||
inputs,
|
||||
return_outputs = return_outputs,
|
||||
num_items_in_batch = num_items_in_batch,
|
||||
)
|
||||
return (loss, outputs) if return_outputs else loss
|
||||
pass
|
||||
|
||||
def sft_trainer_compute_loss(function_name, function):
|
||||
if function_name != "compute_loss": return function
|
||||
|
||||
function = inspect.getsource(_sft_trainer_compute_loss)
|
||||
function = function.replace("def _sft_trainer_compute_loss", "def compute_loss")
|
||||
function = function.split("\n")
|
||||
function = "\n".join(" "*4+x for x in function)
|
||||
def compute_loss(self, model, inputs, return_outputs = False, num_items_in_batch = None):
|
||||
outputs = super().compute_loss(
|
||||
model,
|
||||
inputs,
|
||||
return_outputs = return_outputs,
|
||||
num_items_in_batch = num_items_in_batch,
|
||||
)
|
||||
return outputs
|
||||
pass
|
||||
|
||||
function = inspect.getsource(compute_loss)
|
||||
return function
|
||||
pass
|
||||
RL_FUNCTIONS["sft_trainer"].append(sft_trainer_compute_loss)
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue