Update _utils.py
This commit is contained in:
parent
11e8e19071
commit
b4ef870280
1 changed files with 7 additions and 5 deletions
|
|
@ -1175,17 +1175,19 @@ LOGITS_ERROR_STRING = \
|
|||
"os.environ['UNSLOTH_RETURN_LOGITS'] = '1'\n"\
|
||||
"... trainer.train() ..."
|
||||
|
||||
def raise_logits_error(*args, **kwargs): print(LOGITS_ERROR_STRING)
|
||||
def raise_logits_error1(*args, **kwargs): print(1000)
|
||||
def raise_logits_error2(*args, **kwargs): print(2000)
|
||||
class EmptyLogits:
|
||||
def __init__(self): return
|
||||
__getitem__ = raise_logits_error
|
||||
__getattr__ = raise_logits_error
|
||||
__getitem__ = raise_logits_error1
|
||||
__getattr__ = raise_logits_error2
|
||||
def __repr__(self): return LOGITS_ERROR_STRING
|
||||
def __str__ (self): return LOGITS_ERROR_STRING
|
||||
pass
|
||||
EMPTY_LOGITS = EmptyLogits()
|
||||
functions = dir(torch.Tensor)
|
||||
for function in functions:
|
||||
try: exec(f"EMPTY_LOGITS.{function} = raise_logits_error", globals(), locals())
|
||||
for j, function in enumerate(functions):
|
||||
exec(f"def raise_{j}: print({j})", globals(), locals())
|
||||
try: exec(f"EMPTY_LOGITS.{function} = raise_{j}", globals(), locals())
|
||||
except: continue
|
||||
pass
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue