Update vision.py
This commit is contained in:
parent
0f20d665bd
commit
369ce004df
1 changed files with 4 additions and 10 deletions
|
|
@ -81,7 +81,6 @@ def unsloth_base_fast_generate(
|
|||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
print(args, kwargs)
|
||||
if len(args) != 0:
|
||||
x = args[0]
|
||||
elif "input_ids" in kwargs:
|
||||
|
|
@ -544,15 +543,10 @@ class FastBaseModel:
|
|||
model.for_inference = functools.partial(FastBaseModel.for_inference, model)
|
||||
|
||||
# Patch generate
|
||||
# if model.generate.__name__ != "unsloth_base_fast_generate":
|
||||
# # Check for internal old_generates
|
||||
# m = model
|
||||
# while hasattr(m, "model"):
|
||||
# if hasattr(m, "_old_generate"):
|
||||
|
||||
# model._old_generate = model.generate
|
||||
# unsloth_base_fast_generate.__doc__ = model._old_generate.__doc__
|
||||
# model.generate = types.MethodType(unsloth_base_fast_generate, model)
|
||||
if model.generate.__name__ != "unsloth_base_fast_generate":
|
||||
model._old_generate = model.generate
|
||||
unsloth_base_fast_generate.__doc__ = model._old_generate.__doc__
|
||||
model.generate = types.MethodType(unsloth_base_fast_generate, model)
|
||||
return model
|
||||
pass
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue