# Fine-tune Llama 2 in Google Colab
> 🗣️ Large Language Model Course

❤️ Created by [@maximelabonne](https://twitter.com/maximelabonne), based on Younes Belkada's [GitHub Gist](https://gist.github.com/younesbelkada/9f7f75c94bdc1981c8ca5cc937d4a4da). Special thanks to Tolga HOŞGÖR for his solution to empty the VRAM.

This notebook runs on a T4 GPU. (Last update: 24 Aug 2023)


In [1]:
!pip install -q accelerate==0.21.0 peft==0.4.0 bitsandbytes==0.40.2 transformers==4.31.0 trl==0.4.7

[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m244.2/244.2 kB[0m [31m4.7 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m72.9/72.9 kB[0m [31m7.6 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m92.5/92.5 MB[0m [31m9.1 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m7.4/7.4 MB[0m [31m88.9 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m77.4/77.4 kB[0m [31m10.4 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m7.8/7.8 MB[0m [31m60.5 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m510.5/510.5 kB[0m [31m41.1 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━[0m [32m23.7/23.7 MB[0m [31m42.9 MB/s[0m eta [36m0:00:00[0m
[2K     [90m━━━━━━━━━━━━━━━━━━━━━

In [2]:
from google.colab import drive
drive.mount('/content/drive', force_remount=True)

Mounted at /content/drive


In [3]:
#%cd drive/MyDrive/llama_ft_datasets/task_1/
# %cd task_1
%cd drive/MyDrive/SUPaHOT_data_processing
!ls

/content/drive/MyDrive/SUPaHOT_data_processing
all_resources		eval_old.py  generate_queries.py  llamagrammar.py  oracle.py	  task_1
consolidate_ft_data.py	eval.py      llama2ft_local.py	  meditron_old.py  preprocess.py  task_2
conversation.py		finetune.py  llama2ft.py	  meditron.py	   __pycache__	  task_3
data.tar.gz		ft_datasets  llama2.py		  mock_patients    queries
download_ft_weights.py	ft_model     llama_ft_datasets	  oracle_old.py    results


In [4]:
import os
import torch
from datasets import load_dataset
from transformers import (
    AutoModelForCausalLM,
    AutoTokenizer,
    BitsAndBytesConfig,
    HfArgumentParser,
    TrainingArguments,
    pipeline,
    logging,
)
from peft import LoraConfig, PeftModel
from trl import SFTTrainer

In [5]:


# The model that you want to train from the Hugging Face hub
model_name = "NousResearch/Llama-2-7b-chat-hf"

# The instruction dataset to use
dataset_name = "https://raw.githubusercontent.com/hhekmat/SUPaHOT_data_processing/main/llama_ft_datasets/task_2/task_2_train.jsonl"

# Fine-tuned model name
new_model = "ft_model/post_task_2"

################################################################################
# QLoRA parameters
################################################################################

# LoRA attention dimension
lora_r = 64

# Alpha parameter for LoRA scaling
lora_alpha = 16

# Dropout probability for LoRA layers
lora_dropout = 0.1

################################################################################
# bitsandbytes parameters
################################################################################

# Activate 4-bit precision base model loading
use_4bit = True

# Compute dtype for 4-bit base models
bnb_4bit_compute_dtype = "float16"

# Quantization type (fp4 or nf4)
bnb_4bit_quant_type = "nf4"

# Activate nested quantization for 4-bit base models (double quantization)
use_nested_quant = False

################################################################################
# TrainingArguments parameters
################################################################################

# Output directory where the model predictions and checkpoints will be stored
output_dir = "./results"

# Number of training epochs
num_train_epochs = 1

# Enable fp16/bf16 training (set bf16 to True with an A100)
fp16 = False
bf16 = False

# Batch size per GPU for training
per_device_train_batch_size = 4

# Batch size per GPU for evaluation
per_device_eval_batch_size = 4

# Number of update steps to accumulate the gradients for
gradient_accumulation_steps = 1

# Enable gradient checkpointing
gradient_checkpointing = True

# Maximum gradient normal (gradient clipping)
max_grad_norm = 0.3

# Initial learning rate (AdamW optimizer)
learning_rate = 2e-4

# Weight decay to apply to all layers except bias/LayerNorm weights
weight_decay = 0.001

# Optimizer to use
optim = "paged_adamw_32bit"

# Learning rate schedule
lr_scheduler_type = "cosine"

# Number of training steps (overrides num_train_epochs)
max_steps = -1

# Ratio of steps for a linear warmup (from 0 to learning rate)
warmup_ratio = 0.03

# Group sequences into batches with same length
# Saves memory and speeds up training considerably
group_by_length = True

# Save checkpoint every X updates steps
save_steps = 0

# Log every X updates steps
logging_steps = 25

################################################################################
# SFT parameters
################################################################################

# Maximum sequence length to use
max_seq_length = None

# Pack multiple short examples in the same input sequence to increase efficiency
packing = False

# Load the entire model on the GPU 0
device_map = {"": 0}

In [6]:
# Load dataset (you can process it here)
dataset = load_dataset("json", data_files=dataset_name, split="train")

# Load tokenizer and model with QLoRA configuration
compute_dtype = getattr(torch, bnb_4bit_compute_dtype)

bnb_config = BitsAndBytesConfig(
    load_in_4bit=use_4bit,
    bnb_4bit_quant_type=bnb_4bit_quant_type,
    bnb_4bit_compute_dtype=compute_dtype,
    bnb_4bit_use_double_quant=use_nested_quant,
)

# Check GPU compatibility with bfloat16
if compute_dtype == torch.float16 and use_4bit:
    major, _ = torch.cuda.get_device_capability()
    if major >= 8:
        print("=" * 80)
        print("Your GPU supports bfloat16: accelerate training with bf16=True")
        print("=" * 80)

# Load base model
model = AutoModelForCausalLM.from_pretrained(
    model_name,
    quantization_config=bnb_config,
    device_map=device_map
)
model.config.use_cache = False
model.config.pretraining_tp = 1

# Load LLaMA tokenizer
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right" # Fix weird overflow issue with fp16 training

# Load LoRA configuration
peft_config = LoraConfig(
    lora_alpha=lora_alpha,
    lora_dropout=lora_dropout,
    r=lora_r,
    bias="none",
    task_type="CAUSAL_LM",
)

# Set training parameters
training_arguments = TrainingArguments(
    output_dir=output_dir,
    num_train_epochs=num_train_epochs,
    per_device_train_batch_size=per_device_train_batch_size,
    gradient_accumulation_steps=gradient_accumulation_steps,
    optim=optim,
    save_steps=save_steps,
    logging_steps=logging_steps,
    learning_rate=learning_rate,
    weight_decay=weight_decay,
    fp16=fp16,
    bf16=bf16,
    max_grad_norm=max_grad_norm,
    max_steps=max_steps,
    warmup_ratio=warmup_ratio,
    group_by_length=group_by_length,
    lr_scheduler_type=lr_scheduler_type,
    report_to="tensorboard"
)

# Set supervised fine-tuning parameters
trainer = SFTTrainer(
    model=model,
    train_dataset=dataset,
    peft_config=peft_config,
    dataset_text_field="text",
    max_seq_length=max_seq_length,
    tokenizer=tokenizer,
    args=training_arguments,
    packing=packing,
)

# Train model
trainer.train()

# Save trained model
trainer.model.save_pretrained(new_model)

Downloading data:   0%|          | 0.00/154k [00:00<?, ?B/s]

Generating train split: 0 examples [00:00, ? examples/s]

The secret `HF_TOKEN` does not exist in your Colab secrets.
To authenticate with the Hugging Face Hub, create a token in your settings tab (https://huggingface.co/settings/tokens), set it as secret in your Google Colab and restart your session.
You will be able to reuse this secret in all of your notebooks.
Please note that authentication is recommended but still optional to access public models or datasets.


config.json:   0%|          | 0.00/583 [00:00<?, ?B/s]

model.safetensors.index.json:   0%|          | 0.00/26.8k [00:00<?, ?B/s]

Downloading shards:   0%|          | 0/2 [00:00<?, ?it/s]

model-00001-of-00002.safetensors:   0%|          | 0.00/9.98G [00:00<?, ?B/s]

model-00002-of-00002.safetensors:   0%|          | 0.00/3.50G [00:00<?, ?B/s]

Loading checkpoint shards:   0%|          | 0/2 [00:00<?, ?it/s]

generation_config.json:   0%|          | 0.00/179 [00:00<?, ?B/s]

tokenizer_config.json:   0%|          | 0.00/746 [00:00<?, ?B/s]

tokenizer.model:   0%|          | 0.00/500k [00:00<?, ?B/s]

tokenizer.json:   0%|          | 0.00/1.84M [00:00<?, ?B/s]

added_tokens.json:   0%|          | 0.00/21.0 [00:00<?, ?B/s]

special_tokens_map.json:   0%|          | 0.00/435 [00:00<?, ?B/s]



Map:   0%|          | 0/904 [00:00<?, ? examples/s]

You're using a LlamaTokenizerFast tokenizer. Please note that with a fast tokenizer, using the `__call__` method is faster than using a method to encode the text followed by a call to the `pad` method to get a padded encoding.


Step,Training Loss
25,1.5182
50,0.8852
75,0.694
100,0.6773
125,0.5917
150,0.6218
175,0.5122
200,0.5351
225,0.483


In [None]:
# %load_ext tensorboard
# %tensorboard --logdir results/runs

In [7]:
# Ignore warnings
logging.set_verbosity(logging.CRITICAL)

import json
from datetime import datetime
import os

def parse_fhir_json(file_path):
    global_resource_dict = {}
    relevant_resources = []
    resource_counter = {}
    relevant_resource_types = ['allergyIntolerance', 'Condition', 'Encounter', 'Immunization', 'MedicationRequest', 'Observation', 'Procedure']
    for rt in relevant_resource_types:
        resource_counter[rt] = 0
    with open(file_path, 'r') as f:
        fhir_data = json.load(f)
        if 'entry' in fhir_data:
            for resource in reversed(fhir_data['entry']):
                # print(is_relevant(resource))
                (relevance, rt) = is_relevant(resource)
                if relevance:
                    if resource_counter[rt] < 64:
                        label = extract_display_name_date(resource)
                        relevant_resources.append(label)
                        global_resource_dict[label] = str(resource)
                        resource_counter[rt] += 1
    # print('global resource dict ', global_resource_dict)
    return relevant_resources, global_resource_dict

def is_relevant(resource):
    '''
    RELEVANT RESOURCES
    allergyIntolerances
    + llmConditions (conditions that are active and have proper URL)
    + encounters.uniqueDisplayNames (remove duplicates, only keep most recent instance)
    + immunizations
    + llmMedications (outpatient/active medications)
    + observations.uniqueDisplayNames (remove duplicates, only keep most recent instance)
    + procedures.uniqueDisplayNames (remove duplicates, only keep most recent instance)
    '''
    relevant_resource_types = ['allergyIntolerance', 'Condition', 'Encounter', 'Immunization', 'MedicationRequest', 'Observation', 'Procedure']
    rt = resource['resource']['resourceType']
    if rt not in relevant_resource_types:
        return (False, rt)
    elif rt == 'Condition':
        if resource['resource']['clinicalStatus']['coding'][0]['system'] != 'http://terminology.hl7.org/CodeSystem/condition-clinical' or resource['resource']['clinicalStatus']['coding'][0]['code'] != 'active':
            return (False, rt)
    elif rt == "MedicationRequest":
        if 'medicationCodeableConcept' not in resource['resource'].keys():
            return (False, rt)
    return (True, rt)

def extract_display_name_date(resource):
    '''
    SWIFT CODE FOR FINAL
    var functionCallIdentifier: String {
        resourceType.filter { !$0.isWhitespace }
            + displayName.filter { !$0.isWhitespace }
            + "-"
            + (date.map { FHIRResource.dateFormatter.string(from: $0) } ?? "")
    }
    '''
    def format_date(date):
        input_date = datetime.strptime(date, "%Y-%m-%dT%H:%M:%S%z")
        date = input_date.strftime("%m-%d-%Y")
        return date
    type = resource['resource']['resourceType']
    if type in ['allergyIntolerance', 'Condition']:
        display_name = resource['resource']['code']['text']
        date = format_date(resource['resource']['recordedDate'])
    elif type == 'Encounter':
        display_name = resource['resource']['type'][0]['text']
        date = format_date(resource['resource']['period']['start'])
    elif type == 'Immunization':
        display_name = resource['resource']['vaccineCode']['text']
        date = format_date(resource['resource']['occurrenceDateTime'])
    elif type == 'MedicationRequest':
        display_name = resource['resource']['medicationCodeableConcept']['text']
        date = format_date(resource['resource']['authoredOn'])
    elif type == 'Observation':
        display_name = resource['resource']['code']['text']
        date = format_date(resource['resource']['effectiveDateTime'])
    elif type == 'Procedure':
        display_name = resource['resource']['code']['text']
        date = format_date(resource['resource']['performedPeriod']['start'])
    return (type + ' ' + display_name + ' ' + date)

def populate_global_resources(patient_data_folder):
    resources_data_folder = "./all_resources"
    file_names = os.listdir(patient_data_folder)
    global_resource_dict = {}
    for file_name in file_names:
        if file_name in ('.DS_Store', 'licenses'):
            continue
        file_path = os.path.join(patient_data_folder, file_name)
        relevant_resources, new_resource_dict = parse_fhir_json(file_path)
        global_resource_dict.update(new_resource_dict)
    return global_resource_dict
    '''resources_file_name = file_name[:-5] + 'resources.txt'
        resources_file_path = os.path.join(resources_data_folder, resources_file_name)

        with open(resources_file_path, 'w') as file:
            for item in relevant_resources:
                file.write(item + '\n')
    return global_resource_dict'''


patient_data_folder = "./mock_patients"
resources_data_folder = "./all_resources"
global_resource_dict = populate_global_resources(patient_data_folder)
print(len(global_resource_dict.keys()))


def generate_llama_response(user_prompt, task_prompt):
    prompt = f"<s>[INST] <<SYS>> You are a helpful medical assistant. Users ask you questions about their health care information. You will help and be as concise and clear as possible. {task_prompt} <</SYS>> {user_prompt} [/INST]"
    pipe = pipeline(task="text-generation", model=model, tokenizer=tokenizer, max_new_tokens=700)
    result = pipe(prompt)
    result = result[0]['generated_text']
    idx = result.find('[/INST]')
    return result[idx+7:]

def process_task_2(global_resource_dict):
    base_dir = 'task_1/output/oracle'
    output_dir = 'task_2/output/llama_ft'
    task_2_prompt = "Given an excerpt of a JSON object corresponding to a resource from a patient's FHIR medical records, your job is to provide a brief (1 to 2 sentence) natural language summary of the JSON. Don't explicitly mention that it's a JSON."

    for root, dirs, files in os.walk(base_dir):
        for file in files:
            if file.endswith('.txt'):
                file_path = os.path.join(root, file)
                if file_path.find('test') == -1:
                    continue
                process_file(file_path, root, file, base_dir, output_dir, task_2_prompt, global_resource_dict)
                print('processed')


def process_file(file_path, root, file, base_dir, output_dir, task_2_prompt, global_resource_dict):
    with open(file_path, 'r') as f:
        lines = f.readlines()

    if len(lines) == 0:
        prewritten_response = "No relevant resources were found for this query."
        process_empty_file(root, file, base_dir, output_dir, prewritten_response)
    else:
        for line in lines:
            process_line(line, root, file, base_dir, output_dir, task_2_prompt, global_resource_dict)


def process_empty_file(root, file, base_dir, output_dir, prewritten_response):
    rel_path = os.path.relpath(root, base_dir)
    output_subdir = os.path.join(output_dir, rel_path)
    os.makedirs(output_subdir, exist_ok=True)
    output_file = os.path.join(output_subdir, file)

    with open(output_file, 'a') as f_txt:
        f_txt.write(f"{prewritten_response}\n")

def process_line(line, root, file, base_dir, output_dir, task_2_prompt, global_resource_dict):
    resource_label = line.strip()
    print(resource_label)
    large_resource = global_resource_dict.get(resource_label, '')
    large_resource_str = large_resource #json.dumps(large_resource)
    summary = generate_llama_response("JSON: " + large_resource_str, task_2_prompt)
    summary = ' '.join(summary.split())
    print(summary)

    rel_path = os.path.relpath(root, base_dir)
    output_subdir = os.path.join(output_dir, rel_path)
    os.makedirs(output_subdir, exist_ok=True)
    output_file = os.path.join(output_subdir, file)

    with open(output_file, 'a') as f_txt:
        f_txt.write(f"{summary}\n")

process_task_2(global_resource_dict)

1336
MedicationRequest 24 HR tacrolimus 1 MG Extended Release Oral Tablet 01-05-2023




This record is a Medication Request for a community patient, with a medication code of 24 HR tacrolimus 1 MG Extended Release Oral Tablet, ordered by Dr. Anderson, with a reason of History of renal transplant (situation). The status is marked as stopped.
processed
Procedure Hospice care (regime/therapy) 02-12-2023
This record is about a completed procedure for Hospice care that took place at NORTH RIVER HOSPICE LLC on February 12, 2023. The procedure was performed on a specific patient during a specific encounter.
Procedure Hospice care (regime/therapy) 02-25-2023
This record documents a completed procedure for Hospice care that took place at NORTH RIVER HOSPICE LLC on February 25, 2023. The procedure was performed on a specific patient during a specific encounter.
Procedure Hospice care (regime/therapy) 03-02-2023
This record is about a completed procedure for Hospice care that took place at NORTH RIVER HOSPICE LLC on March 2, 2023. The procedure was performed on a specific patient du

In [None]:
# Empty VRAM
del model
del pipe
del trainer
import gc
gc.collect()
gc.collect()

19965

In [None]:
# Reload model in FP16 and merge it with LoRA weights
base_model = AutoModelForCausalLM.from_pretrained(
    model_name,
    low_cpu_mem_usage=True,
    return_dict=True,
    torch_dtype=torch.float16,
    device_map=device_map,
)
model = PeftModel.from_pretrained(base_model, new_model)
model = model.merge_and_unload()

# Reload tokenizer to save it
tokenizer = AutoTokenizer.from_pretrained(model_name, trust_remote_code=True)
tokenizer.pad_token = tokenizer.eos_token
tokenizer.padding_side = "right"

Loading checkpoint shards:   0%|          | 0/2 [00:00<?, ?it/s]

In [None]:
!huggingface-cli login

model.push_to_hub(new_model, use_temp_dir=False)
tokenizer.push_to_hub(new_model, use_temp_dir=False)


    _|    _|  _|    _|    _|_|_|    _|_|_|  _|_|_|  _|      _|    _|_|_|      _|_|_|_|    _|_|      _|_|_|  _|_|_|_|
    _|    _|  _|    _|  _|        _|          _|    _|_|    _|  _|            _|        _|    _|  _|        _|
    _|_|_|_|  _|    _|  _|  _|_|  _|  _|_|    _|    _|  _|  _|  _|  _|_|      _|_|_|    _|_|_|_|  _|        _|_|_|
    _|    _|  _|    _|  _|    _|  _|    _|    _|    _|    _|_|  _|    _|      _|        _|    _|  _|        _|
    _|    _|    _|_|      _|_|_|    _|_|_|  _|_|_|  _|      _|    _|_|_|      _|        _|    _|    _|_|_|  _|_|_|_|
    
    To login, `huggingface_hub` requires a token generated from https://huggingface.co/settings/tokens .
Token: 
Add token as git credential? (Y/n) n
Token is valid (permission: write).
Your token has been saved to /root/.cache/huggingface/token
Login successful


Upload 2 LFS files:   0%|          | 0/2 [00:00<?, ?it/s]

pytorch_model-00001-of-00002.bin:   0%|          | 0.00/9.98G [00:00<?, ?B/s]

pytorch_model-00002-of-00002.bin:   0%|          | 0.00/3.50G [00:00<?, ?B/s]

CommitInfo(commit_url='https://huggingface.co/mlabonne/llama-2-7b-miniguanaco/commit/c81a32fd0b4d39e252326e639d63e75aa68c9a4a', commit_message='Upload tokenizer', commit_description='', oid='c81a32fd0b4d39e252326e639d63e75aa68c9a4a', pr_url=None, pr_revision=None, pr_num=None)