make unsloth_tiled_mlp a from_pretrained arg (#3655)
* make unsloth_tiled_mlp a from_pretrained arg * adjust patching logic
This commit is contained in:
parent
27ae5c335c
commit
69f7dfdb90
1 changed files with 11 additions and 4 deletions
|
|
@ -146,6 +146,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
disable_log_stats = True,
|
||||
qat_scheme = None,
|
||||
load_in_fp8 = False, # fp8 LoRA (True, False, 'block')
|
||||
unsloth_tiled_mlp = False,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -512,6 +513,7 @@ class FastLanguageModel(FastLlamaModel):
|
|||
disable_log_stats = disable_log_stats,
|
||||
qat_scheme = qat_scheme,
|
||||
load_in_fp8 = load_in_fp8,
|
||||
unsloth_tiled_mlp = unsloth_tiled_mlp,
|
||||
*args,
|
||||
**kwargs,
|
||||
)
|
||||
|
|
@ -609,8 +611,10 @@ class FastLanguageModel(FastLlamaModel):
|
|||
|
||||
# Patch Tiled MLP
|
||||
# to turn on set UNSLOTH_TILED_MLP to "arctic", "target", or "target:{GB}""
|
||||
patch_tiled_mlp_choice = os.environ.get("UNSLOTH_TILED_MLP", "0")
|
||||
if patch_tiled_mlp_choice != "0":
|
||||
patch_tiled_mlp_choice = os.environ.get(
|
||||
"UNSLOTH_TILED_MLP", "arctic" if unsloth_tiled_mlp else "0"
|
||||
)
|
||||
if patch_tiled_mlp_choice != "0" or unsloth_tiled_mlp:
|
||||
patch_tiled_mlp(model, patch_options_str = patch_tiled_mlp_choice)
|
||||
|
||||
return model, tokenizer
|
||||
|
|
@ -674,6 +678,7 @@ class FastModel(FastBaseModel):
|
|||
disable_log_stats = True,
|
||||
qat_scheme = None,
|
||||
load_in_fp8 = False, # fp8 LoRA (True, False, 'block')
|
||||
unsloth_tiled_mlp = False,
|
||||
*args,
|
||||
**kwargs,
|
||||
):
|
||||
|
|
@ -1240,8 +1245,10 @@ class FastModel(FastBaseModel):
|
|||
|
||||
# Patch Tiled MLP
|
||||
# to turn on set UNSLOTH_TILED_MLP to "arctic", "target", or "target:{GB}""
|
||||
patch_tiled_mlp_choice = os.environ.get("UNSLOTH_TILED_MLP", "0")
|
||||
if patch_tiled_mlp_choice != "0":
|
||||
patch_tiled_mlp_choice = os.environ.get(
|
||||
"UNSLOTH_TILED_MLP", "arctic" if unsloth_tiled_mlp else "0"
|
||||
)
|
||||
if patch_tiled_mlp_choice != "0" or unsloth_tiled_mlp:
|
||||
patch_tiled_mlp(model, patch_options_str = patch_tiled_mlp_choice)
|
||||
|
||||
return model, tokenizer
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue