Update _utils.py
This commit is contained in:
parent
fcb40cc366
commit
532db6c564
1 changed files with 4 additions and 4 deletions
|
|
@ -1245,15 +1245,15 @@ pass
|
|||
# os.environ['UNSLOTH_RETURN_LOGITS'] = '1'
|
||||
LOGITS_ERROR_STRING = \
|
||||
"Unsloth: Logits are empty from 2024.11 onwards. To get raw logits again, please "\
|
||||
'set the environment variable `UNSLOTH_RETURN_LOGITS` to `"1" BEFORE starting to train ie before `trainer.train()`. For example:\n\n'\
|
||||
"import os\n"\
|
||||
'set the environment variable `UNSLOTH_RETURN_LOGITS` to `"1" BEFORE starting to train ie before `trainer.train()`. For example:\n'\
|
||||
"```\nimport os\n"\
|
||||
"os.environ['UNSLOTH_RETURN_LOGITS'] = '1'\n"\
|
||||
"... trainer.train() ...\n"\
|
||||
"trainer.train()\n```\n"\
|
||||
"No need to restart your console - just add `os.environ['UNSLOTH_RETURN_LOGITS'] = '1'` before trainer.train() and re-run the cell!"
|
||||
|
||||
def raise_logits_error(*args, **kwargs): raise NotImplementedError(LOGITS_ERROR_STRING)
|
||||
def return_none(*args, **kwargs): return None
|
||||
class EmptyLogits:
|
||||
class EmptyLogits(torch.Tensor):
|
||||
def __init__(self): return
|
||||
def raise_getattr_error(self, attr): return return_none if attr == "to" else raise_logits_error
|
||||
__getitem__ = raise_logits_error
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue