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