distributed data parallelism with torchrun

This commit is contained in:
kooshi
2023-03-24 23:56:06 -05:00
parent 2bc64597aa
commit 8e471516b8
4 changed files with 20 additions and 7 deletions

View File

@@ -197,7 +197,7 @@ def model_to_float(model):
print('Converted as Float.')
def load_llama_model_4bit_low_ram(config_path, model_path, half=False):
def load_llama_model_4bit_low_ram(config_path, model_path, half=False, device_map="auto"):
import transformers
import accelerate
from transformers import LlamaConfig, LlamaForCausalLM, LlamaTokenizer
@@ -222,7 +222,7 @@ def load_llama_model_4bit_low_ram(config_path, model_path, half=False):
model = accelerate.load_checkpoint_and_dispatch(
model=model,
checkpoint=model_path,
device_map='auto',
device_map=device_map,
no_split_module_classes=["LlamaDecoderLayer"]
)