| 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 |
|
|