mirror of
https://github.com/wassname/Open-Assistant.git
synced 2026-08-19 11:50:35 +08:00
better
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
from transformers import AutoModelForCausalLM
|
||||
import transformers
|
||||
|
||||
# from .gptj import get_model as get_gptj_model
|
||||
|
||||
@@ -25,9 +25,12 @@ def freeze_top_n_layers(model, target_layers):
|
||||
return model
|
||||
|
||||
|
||||
def get_specific_model(model_name, cache_dir, quantization):
|
||||
return AutoModelForCausalLM.from_pretrained(model_name, cache_dir=cache_dir)
|
||||
# if "gpt-j" in model_name.lower():
|
||||
# return get_gptj_model(model_name, cache_dir, quantization)
|
||||
# else:
|
||||
# return AutoModelForCausalLM.from_pretrained(model_name, cache_dir=cache_dir)
|
||||
def get_specific_model(model_name, cache_dir, quantization, seq2seqmodel):
|
||||
# encoder-decoder support for Flan-T5 like models
|
||||
# for now, we can use an argument but in the future,
|
||||
# we can automate this
|
||||
if seq2seqmodel:
|
||||
model = transformers.AutoModelForSeq2SeqLM.from_pretrained(model_name, cache_dir=cache_dir)
|
||||
else:
|
||||
model = transformers.AutoModelForCausalLM.from_pretrained(model_name, cache_dir=cache_dir)
|
||||
return model
|
||||
|
||||
Reference in New Issue
Block a user