TrinityX / models /base_model.py
Gautam Kashyap
Upload folder using huggingface_hub
e38f140 verified
Raw
History Blame Contribute Delete
1.42 kB
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