File size: 1,416 Bytes
e38f140 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 | import os
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
def load_tokenizer(model_name: str = "meta-llama/Llama-2-7b-chat-hf", hf_token: str = None):
token = hf_token or os.environ.get("HF_TOKEN")
tokenizer = AutoTokenizer.from_pretrained(model_name, token=token, use_fast=False)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
tokenizer.pad_token_id = tokenizer.eos_token_id
tokenizer.padding_side = "right"
return tokenizer
def load_base_model(
model_name: str = "meta-llama/Llama-2-7b-chat-hf",
precision: str = "bfloat16",
device_map: str = None,
hf_token: str = None,
freeze: bool = True,
):
token = hf_token or os.environ.get("HF_TOKEN")
dtype_map = {
"fp32": torch.float32,
"bfloat16": torch.bfloat16,
"bf16": torch.bfloat16,
"fp16": torch.float16,
"float16": torch.float16,
}
torch_dtype = dtype_map.get(precision, torch.bfloat16)
model = AutoModelForCausalLM.from_pretrained(
model_name, dtype=torch_dtype, device_map=device_map, token=token)
if freeze:
model.eval()
for p in model.parameters():
p.requires_grad = False
return model
def get_llama_hidden_size(model) -> int:
return model.config.hidden_size
def get_num_layers(model) -> int:
return model.config.num_hidden_layers
|